Files
AI/참고/guardrails-main/guardrails/applications/text2sql.py
2026-05-12 19:40:31 +09:00

220 lines
7.3 KiB
Python

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