참고소스 수정본
This commit is contained in:
219
참고/guardrails-main/guardrails/applications/text2sql.py
Normal file
219
참고/guardrails-main/guardrails/applications/text2sql.py
Normal file
@@ -0,0 +1,219 @@
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import openai
|
||||
from string import Template
|
||||
from typing import Callable, Dict, Optional, Type, cast
|
||||
|
||||
from guardrails.classes import ValidationOutcome
|
||||
from guardrails.document_store import DocumentStoreBase, EphemeralDocumentStore
|
||||
from guardrails.embedding import EmbeddingBase, OpenAIEmbedding
|
||||
from guardrails.guard import Guard
|
||||
from guardrails.utils.sql_utils import create_sql_driver
|
||||
from guardrails.vectordb import Faiss, VectorDBBase
|
||||
|
||||
REASK_PROMPT = """
|
||||
You are a data scientist whose job is to write SQL queries.
|
||||
|
||||
${gr.complete_json_suffix_v2}
|
||||
|
||||
Here's schema about the database that you can use to generate the SQL query.
|
||||
Try to avoid using joins if the data can be retrieved from the same table.
|
||||
|
||||
${db_info}
|
||||
|
||||
I will give you a list of examples.
|
||||
|
||||
${examples}
|
||||
|
||||
I want to create a query for the following instruction:
|
||||
|
||||
${nl_instruction}
|
||||
|
||||
For this instruction, I was given the following JSON, which has some incorrect values.
|
||||
|
||||
${previous_response}
|
||||
|
||||
Help me correct the incorrect values based on the given error messages.
|
||||
"""
|
||||
|
||||
|
||||
EXAMPLE_BOILERPLATE = """
|
||||
I will give you a list of examples. Write a SQL query similar to the examples below:
|
||||
"""
|
||||
|
||||
|
||||
def example_formatter(
|
||||
input: str, output: str, output_schema: Optional[Callable] = None
|
||||
) -> str:
|
||||
if output_schema is not None:
|
||||
output = output_schema(output)
|
||||
|
||||
example = "\nINSTRUCTIONS:\n============\n"
|
||||
example += f"{input}\n\n"
|
||||
|
||||
example += "SQL QUERY:\n================\n"
|
||||
example += f"{output}\n\n"
|
||||
|
||||
return example
|
||||
|
||||
|
||||
class Text2Sql:
|
||||
def __init__(
|
||||
self,
|
||||
conn_str: str,
|
||||
schema_file: Optional[str] = None,
|
||||
examples: Optional[Dict] = None,
|
||||
embedding: Type[EmbeddingBase] = OpenAIEmbedding,
|
||||
vector_db: Type[VectorDBBase] = Faiss,
|
||||
document_store: Type[DocumentStoreBase] = EphemeralDocumentStore,
|
||||
rail_spec: Optional[str] = None,
|
||||
rail_params: Optional[Dict] = None,
|
||||
example_formatter: Callable = example_formatter,
|
||||
reask_messages: list[Dict[str, str]] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": REASK_PROMPT,
|
||||
}
|
||||
],
|
||||
llm_api: Optional[Callable] = None,
|
||||
llm_api_kwargs: Optional[Dict] = None,
|
||||
num_relevant_examples: int = 2,
|
||||
):
|
||||
"""Initialize the text2sql application.
|
||||
|
||||
Args:
|
||||
conn_str: Connection string to the database.
|
||||
schema_file: Path to the schema file. Defaults to None.
|
||||
examples: Examples to add to the document store. Defaults to None.
|
||||
embedding: Embedding to use for document store. Defaults to OpenAIEmbedding.
|
||||
vector_db: Vector database to use for the document store. Defaults to Faiss.
|
||||
document_store: Document store to use. Defaults to EphemeralDocumentStore.
|
||||
rail_spec: Path to the rail specification. Defaults to "text2sql.rail".
|
||||
example_formatter: Fn to format examples. Defaults to example_formatter.
|
||||
reask_prompt: Prompt to use for reasking. Defaults to REASK_PROMPT.
|
||||
"""
|
||||
if llm_api is None:
|
||||
llm_api = openai.completions.create
|
||||
|
||||
self.example_formatter = example_formatter
|
||||
self.llm_api = llm_api
|
||||
self.llm_api_kwargs = llm_api_kwargs or {"max_tokens": 512}
|
||||
|
||||
# Initialize the SQL driver.
|
||||
self.sql_driver = create_sql_driver(conn=conn_str, schema_file=schema_file)
|
||||
self.sql_schema = self.sql_driver.get_schema()
|
||||
|
||||
# Number of relevant examples to use for the LLM.
|
||||
self.num_relevant_examples = num_relevant_examples
|
||||
|
||||
# Initialize the Guard class.
|
||||
self.guard = self._init_guard(
|
||||
conn_str,
|
||||
schema_file,
|
||||
rail_spec,
|
||||
rail_params,
|
||||
reask_messages,
|
||||
)
|
||||
|
||||
# Initialize the document store.
|
||||
self.store = self._create_docstore_with_examples(
|
||||
examples, embedding, vector_db, document_store
|
||||
)
|
||||
|
||||
def _init_guard(
|
||||
self,
|
||||
conn_str: str,
|
||||
schema_file: Optional[str] = None,
|
||||
rail_spec: Optional[str] = None,
|
||||
rail_params: Optional[Dict] = None,
|
||||
reask_messages: list[Dict[str, str]] = [
|
||||
{
|
||||
"role": "user",
|
||||
"content": REASK_PROMPT,
|
||||
}
|
||||
],
|
||||
):
|
||||
# Initialize the Guard class
|
||||
if rail_spec is None:
|
||||
rail_spec = os.path.join(os.path.dirname(__file__), "text2sql.rail")
|
||||
rail_params = {"conn_str": conn_str, "schema_file": schema_file}
|
||||
if schema_file is None:
|
||||
rail_params["schema_file"] = ""
|
||||
|
||||
# Load the rail specification.
|
||||
with open(rail_spec, "r") as f:
|
||||
rail_spec_str = f.read()
|
||||
|
||||
# Substitute the parameters in the rail specification.
|
||||
if rail_params is not None:
|
||||
rail_spec_str = Template(rail_spec_str).safe_substitute(**rail_params)
|
||||
|
||||
guard = Guard.for_rail_string(rail_spec_str)
|
||||
guard._exec_opts.reask_messages = reask_messages
|
||||
|
||||
return guard
|
||||
|
||||
def _create_docstore_with_examples(
|
||||
self,
|
||||
examples: Optional[Dict],
|
||||
embedding: Type[EmbeddingBase],
|
||||
vector_db: Type[VectorDBBase],
|
||||
document_store: Type[DocumentStoreBase],
|
||||
) -> Optional[DocumentStoreBase]:
|
||||
if examples is None:
|
||||
return None
|
||||
|
||||
"""Add examples to the document store."""
|
||||
e = embedding()
|
||||
if vector_db == Faiss:
|
||||
db = Faiss.new_flat_l2_index(e.output_dim, embedder=e)
|
||||
else:
|
||||
raise NotImplementedError(f"VectorDB {vector_db} is not implemented.")
|
||||
store = document_store(db)
|
||||
store.add_texts(
|
||||
{example["question"]: {"ctx": example["query"]} for example in examples}
|
||||
)
|
||||
return store
|
||||
|
||||
@staticmethod
|
||||
def output_schema_formatter(output) -> str:
|
||||
return json.dumps({"generated_sql": output}, indent=4)
|
||||
|
||||
def __call__(self, text: str) -> Optional[str]:
|
||||
"""Run text2sql on a text query and return the SQL query."""
|
||||
|
||||
if self.store is not None:
|
||||
similar_examples = self.store.search(text, self.num_relevant_examples)
|
||||
similar_examples_prompt = "\n".join(
|
||||
self.example_formatter(example.text, example.metadata["ctx"])
|
||||
for example in similar_examples
|
||||
)
|
||||
else:
|
||||
similar_examples_prompt = ""
|
||||
|
||||
if asyncio.iscoroutinefunction(self.llm_api):
|
||||
raise ValueError(
|
||||
"Async API is not supported in Text2SQL application. "
|
||||
"Please use a synchronous API."
|
||||
)
|
||||
else:
|
||||
if self.llm_api is None:
|
||||
return None
|
||||
try:
|
||||
response = self.guard(
|
||||
self.llm_api,
|
||||
prompt_params={
|
||||
"nl_instruction": text,
|
||||
"examples": similar_examples_prompt,
|
||||
"db_info": str(self.sql_schema),
|
||||
},
|
||||
**self.llm_api_kwargs,
|
||||
)
|
||||
response = cast(ValidationOutcome, response)
|
||||
validated_output: Dict = cast(Dict, response.validated_output)
|
||||
output = validated_output["generated_sql"]
|
||||
except TypeError:
|
||||
output = None
|
||||
|
||||
return output
|
||||
Reference in New Issue
Block a user