Feature(ai service): refine ask/ask details pipelines (#64)

* allow visualization using one file

* remove default argparse option value

* resolve conflict

* resolve conflict

* resolve conflict

* update setup instructions

* resolve conflict

* resolve conflicts

* resolve conflict

* resolve conflict

* resolve conflict

* add anthropic model pricing

* allow change model for anthropic

* resolve conflict

* resolve conflict

* resovle conflict

* resovle conflict

* resolve conflict

* resolve conflict

* add container:ubuntu

* undo

* resolve conflict

* resolve conflict

* resolve conflict

* resolve conflict

* resolve conflict

* update

* resolve conflict

* resolve conflicts

* resolve conflicts

* resolve conflicts

* resolve conflicts

* resolve conflicts

* resolve conflicts

* resolve conflicts

* resolve conflicts

* remove unused files

* resolve conflict

* resolve conflict

* generate ddl from mdl for indexing

* resolve conflict

* resolve conflicts

* eval ask with semantic description (#45)

* update make eval-ask

* add semantic description to mdl before eval

* refine

* add model semantic generation

* simplify make command

* print timestamp

* minor update

---------

Co-authored-by: qa <qa@qadeMacBook-Pro.local>
Co-authored-by: ChihYu Yeh <chihyu.jimmy.yeh@gmail.com>
Co-authored-by: Aster Sun <aster.sun@cannerdata.com>

* remove ;

* simplify prompt

* revert topk

* update prompt

* make default top_k for retriever 10

* generate multiple results for ask pipeline

* resolve conflict

* add followup_generation_pipeline

* add logging for ask api results

* fix generate_mdl issue

* fix ask_details error message

* add extra standard package for uvicorn

* refine ask

* fix postprocessor for ask details

* fix preview data

* fix test

* update prompts

* refine prompts for ask and ask details

* fix prompt

* refine prompt

* remove slash at the end of api endpoints

* refine prompt

* add custom semantic run option

* resolve conflict

* fix bugs for tests and eval and restructure eval-ask

* fix trainling slash and add logging

* remove unused code

* restructure eval ask pipeline outputs

* add comments to the eval command

* change generation component argument name

* add try/except to handle other errors that might happen

* change WREN_AI_SERVICE_VERSION to nightly

* fix eval ask_details errors

* add error message to indexing

---------

Co-authored-by: imAsterSun <61279528+imAsterSun@users.noreply.github.com>
Co-authored-by: qa <qa@qadeMacBook-Pro.local>
Co-authored-by: Aster Sun <aster.sun@cannerdata.com>
Co-authored-by: Pao Sheng <paooap.oappao@gmail.com>
This commit is contained in:
Chih-Yu Yeh
2024-04-08 15:44:15 +08:00
committed by GitHub
co-authored by qa Aster Sun imAsterSun Pao Sheng
parent b677dd1f77
commit 9507838991
37 changed files with 1597 additions and 1521 deletions
+1 -1
View File
@@ -9,7 +9,7 @@ WREN_AI_SERVICE_PORT=5555
# version
# CHANGE THIS TO THE LATEST VERSION
WREN_ENGINE_VERSION=nightly
WREN_AI_SERVICE_VERSION=dev
WREN_AI_SERVICE_VERSION=nightly
WREN_UI_VERSION=0.1.0
WREN_BOOTSTRAP_VERSION=0.1.0
+1 -4
View File
@@ -4,11 +4,8 @@ WREN_AI_SERVICE_PORT=5555
# app related
QDRANT_HOST=localhost
OPENAI_API_KEY=
LANGFUSE_PUBLIC_KEY=
LANGFUSE_SECRET_KEY=
ENABLE_TRACE=
WREN_ENGINE_ENDPOINT=http://localhost:8080
OPENAI_API_KEY=
# evaluation related
DATASET_NAME=book_2
+11 -13
View File
@@ -23,15 +23,6 @@ run-qdrant:
stop-qdrant:
docker stop qdrant && docker rm qdrant
# present the evaluation result on the streamlit app
# example: make streamlit pipeline=src/eval/streamlit_app.py
streamlit:
poetry run streamlit run $(pipeline)
# example: make eval pipeline=ask_details
eval:
poetry run python -m src.eval.$(pipeline) $(args)
run-wren-engine:
docker compose -f ./src/eval/wren-engine/docker-compose.yml --env-file ./src/eval/wren-engine/.env up -d
@@ -50,14 +41,21 @@ stop-all:
make stop-qdrant && \
make stop-wren-engine
eval-ask:
# present the evaluation result on the streamlit app
# example: make streamlit pipeline=ask_details
streamlit:
poetry run streamlit run src/eval/${pipeline}/streamlit_app.py
# example: make eval pipeline=ask_details
# example: make eval pipeline=ask args="--help" to check all available arguments
eval:
make run-all && \
poetry run python -m src.eval.ask --eval-from-scratch --eval-after-prediction && \
poetry run python -m src.eval.$(pipeline) $(args)
make stop-all
test:
poetry run python -m src.prepare_mdl_json --dataset_name book_2 && \
make run-qdrant && \
make run-wren-engine && \
poetry run pytest -s && \
make stop-all
poetry run pytest -s $(args) && \
make stop-all
+1
View File
@@ -22,6 +22,7 @@
## Pipeline Evaluation(for development)
- install `psql`
- fill in environment variables: `.env.dev` in the src folder and `config.properties` in the src/eval/wren-engine/etc folder
- start the docker service
- run qdrant and wren-engine docker containers: `make run-all`
+46 -51
View File
@@ -336,39 +336,31 @@ def generate_mdl_json(
}
)
else:
should_add_column = True
if "PRIMARY KEY" in part or "primary key" in part:
if "(" not in part and ")" not in part:
primary_key = part.strip().split(" ")[0]
part = (
part.replace("PRIMARY KEY", "")
.replace("primary key", "")
.strip()
)
else:
should_add_column = False
pattern = r'\("(.*?)"\)'
if matches := re.findall(pattern, part):
primary_key = matches[0]
primary_key = part.strip().split(" ")[0]
part = (
part.replace("PRIMARY KEY", "")
.replace("primary key", "")
.strip()
)
# Splitting the column name and type
if should_add_column:
column_def = _parse_column_definition(part.strip())
column_def = _parse_column_definition(part.strip())
columns.append(
{
"name": column_def["name"].replace('"', ""),
"type": _get_appropriat_column_type(column_def["type"]),
"notNull": column_def[
"not_null"
], # Assuming notNull is False by default as not specified in the string
"isCalculated": False, # Assuming isCalculated is False by default
"expression": column_def["name"].replace(
'"', ""
), # Assuming expression is the column name itself
"properties": {},
}
)
columns.append(
{
"name": column_def["name"].replace('"', ""),
"type": _get_appropriat_column_type(column_def["type"]),
"notNull": column_def[
"not_null"
], # Assuming notNull is False by default as not specified in the string
"isCalculated": False, # Assuming isCalculated is False by default
"expression": column_def["name"].replace(
'"', ""
), # Assuming expression is the column name itself
"properties": {},
}
)
if relationships:
for relationship in relationships:
@@ -626,22 +618,18 @@ def show_asks_details_results():
for i, step in enumerate(st.session_state["asks_details_result"]["steps"]):
st.markdown(f"#### Step {i + 1}")
st.markdown(step["summary"])
if i != len(st.session_state["asks_details_result"]["steps"]) - 1:
st.code(
body=step["sql"],
language="sql",
)
sqls_with_cte.append(
"WITH " + step["cte_name"] + " AS (" + step["sql"] + ")"
)
sqls.append(step["sql"])
else:
last_step_sql = "\n".join(sqls_with_cte) + "\n\n" + step["sql"]
sqls.append(last_step_sql)
st.code(
body=last_step_sql,
language="sql",
)
sql = ""
if sqls_with_cte:
sql += "WITH " + ",\n".join(sqls_with_cte) + "\n\n"
sql += step["sql"]
sqls.append(sql)
st.code(
body=sql,
language="sql",
)
sqls_with_cte.append(f"{step['cte_name']} AS ( {step['sql']} )")
st.button(
label="Preview Data",
@@ -679,7 +667,7 @@ def generate_mdl_metadata(mdl_model_json: dict):
st.toast(f'Generating MDL metadata for model {mdl_model_json['name']}', icon="⏳")
generate_mdl_metadata_response = requests.post(
f"{WREN_AI_SERVICE_BASE_URL}/v1/semantics-descriptions/",
f"{WREN_AI_SERVICE_BASE_URL}/v1/semantics-descriptions",
json={
"mdl": mdl_model_json,
"model": mdl_model_json["name"],
@@ -708,7 +696,7 @@ def generate_mdl_metadata(mdl_model_json: dict):
def prepare_semantics(mdl_json: dict):
semantics_preparation_response = requests.post(
f"{WREN_AI_SERVICE_BASE_URL}/v1/semantics-preparations/",
f"{WREN_AI_SERVICE_BASE_URL}/v1/semantics-preparations",
json={
"mdl": json.dumps(mdl_json),
"id": st.session_state["deployment_id"],
@@ -750,7 +738,7 @@ def prepare_semantics(mdl_json: dict):
def ask(query: str, query_history: Optional[dict] = None):
st.session_state["query"] = query
asks_response = requests.post(
f"{WREN_AI_SERVICE_BASE_URL}/v1/asks/",
f"{WREN_AI_SERVICE_BASE_URL}/v1/asks",
json={
"query": query,
"id": st.session_state["deployment_id"],
@@ -768,7 +756,7 @@ def ask(query: str, query_history: Optional[dict] = None):
and asks_status != "stopped"
):
asks_status_response = requests.get(
f"{WREN_AI_SERVICE_BASE_URL}/v1/asks/{query_id}/result/"
f"{WREN_AI_SERVICE_BASE_URL}/v1/asks/{query_id}/result"
)
assert asks_status_response.status_code == 200
asks_status = asks_status_response.json()["status"]
@@ -786,7 +774,7 @@ def ask(query: str, query_history: Optional[dict] = None):
def ask_details():
asks_details_response = requests.post(
f"{WREN_AI_SERVICE_BASE_URL}/v1/ask-details/",
f"{WREN_AI_SERVICE_BASE_URL}/v1/ask-details",
json={
"query": st.session_state["chosen_query_result"]["query"],
"sql": st.session_state["chosen_query_result"]["sql"],
@@ -798,7 +786,9 @@ def ask_details():
query_id = asks_details_response.json()["query_id"]
asks_details_status = None
while not asks_details_status or asks_details_status != "finished":
while (
asks_details_status != "finished" and asks_details_status != "failed"
) or not asks_details_status:
asks_details_status_response = requests.get(
f"{WREN_AI_SERVICE_BASE_URL}/v1/ask-details/{query_id}/result/"
)
@@ -811,3 +801,8 @@ def ask_details():
st.session_state["asks_details_result"] = asks_details_status_response.json()[
"response"
]
elif asks_details_status == "failed":
st.error(
f'An error occurred while processing the query: {asks_details_status_response.json()['error']}',
icon="🚨",
)
+1 -2
View File
@@ -8,7 +8,7 @@ readme = "README.md"
[tool.poetry.dependencies]
python = "^3.12"
fastapi = "^0.109.2"
uvicorn = "^0.27.1"
uvicorn = {extras = ["standard"], version = "^0.29.0"}
python-dotenv = "^1.0.1"
haystack-ai = "^2.0.0"
openai = "^1.14.0"
@@ -16,7 +16,6 @@ qdrant-haystack = "^3.0.0"
backoff = "^2.2.1"
tqdm = "^4.66.2"
numpy = "^1.26.4"
langfuse = "^2.19.1"
[tool.poetry.group.dev.dependencies]
pytest = "^8.0.0"
+7 -17
View File
@@ -14,25 +14,15 @@ class BasicPipeline(metaclass=ABCMeta):
def run(self, *args, **kwargs) -> Dict[str, Any]:
...
def save(self, with_trace: bool = False, suffix: str = None) -> Path:
def save(self, suffix: str = None) -> Path:
if suffix:
if with_trace:
file_path = Path(
f"./outputs/{self.__class__.__name__.lower()}_pipeline_with_trace_{suffix}.yaml"
)
else:
file_path = Path(
f"./outputs/{self.__class__.__name__.lower()}_pipeline_{suffix}.yaml"
)
file_path = Path(
f"./outputs/{self.__class__.__name__.lower()}_pipeline_{suffix}.yaml"
)
else:
if with_trace:
file_path = Path(
f"./outputs/{self.__class__.__name__.lower()}_pipeline_with_trace.yaml"
)
else:
file_path = Path(
f"./outputs/{self.__class__.__name__.lower()}_pipeline.yaml"
)
file_path = Path(
f"./outputs/{self.__class__.__name__.lower()}_pipeline.yaml"
)
with open(file_path, "w") as file:
self._pipe.dump(file)
@@ -2,7 +2,6 @@ import argparse
import json
import os
import time
import uuid
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from pathlib import Path
@@ -18,7 +17,12 @@ from src.pipelines.ask.generation_pipeline import Generation
from src.pipelines.ask.indexing_pipeline import Indexing
from src.pipelines.ask.retrieval_pipeline import Retrieval
from src.pipelines.ask.sql_correction_pipeline import SQLCorrection
from src.pipelines.semantics import description
from src.utils import load_env_vars
from src.web.v1.services.semantics import (
GenerateDescriptionRequest,
SemanticsService,
)
# from .eval_pipeline import Evaluation
from .utils import (
@@ -30,19 +34,16 @@ from .utils import (
load_env_vars()
if with_trace := os.getenv("ENABLE_TRACE", default=False):
from src.pipelines.trace import (
langfuse,
)
def process_item(query: str, user_id: Optional[str] = None) -> Dict[str, Any]:
def process_item(query: str, no_db_schema: Optional[bool]) -> Dict[str, Any]:
retrieval_start = time.perf_counter()
retrieval_result = retrieval_pipeline.run(
query,
user_id=user_id,
)
documents = retrieval_result["post_processor"]["documents"]
if not no_db_schema:
retrieval_result = retrieval_pipeline.run(
query,
)
documents = retrieval_result["post_processor"]["documents"]
else:
documents = []
retrieval_end = time.perf_counter()
valid_generation_results = []
@@ -64,7 +65,6 @@ def process_item(query: str, user_id: Optional[str] = None) -> Dict[str, Any]:
text_to_sql_generation_results = generation_pipeline.run(
query,
contexts=documents,
user_id=user_id,
)
text_to_sql_generation_end = time.perf_counter()
text_to_sql_generation_time_cost = (
@@ -193,7 +193,7 @@ def eval(prediction_results_file: Path, dataset_name: str, ground_truths: list[d
timestamp = prediction_results_file.stem.split("_")[-1]
with open(f"./outputs/{dataset_name}_eval_results_{timestamp}.json", "w") as f:
with open(f"./outputs/ask/{dataset_name}_eval_results_{timestamp}.json", "w") as f:
json.dump(eval_results, f, indent=2)
@@ -206,54 +206,151 @@ if __name__ == "__main__":
parser.add_argument(
"--input-file",
type=str,
default=get_latest_prediction_outputs_file(Path("./outputs"), DATASET_NAME),
default=get_latest_prediction_outputs_file(Path("./outputs/ask"), DATASET_NAME),
help="Path to the prediction results file. If not provided, the latest prediction results file will be used. The file should be located in the outputs folder in the root directory of the project.",
)
parser.add_argument(
"--eval-after-prediction",
action=argparse.BooleanOptionalAction,
help="Whether to run the evaluation after making predictions. Default is True.",
help="Run the evaluation after making predictions. Default is False.",
)
parser.add_argument(
"--eval-from-scratch",
action=argparse.BooleanOptionalAction,
help="Whether to run the evaluation from scratch. Default is False.",
help="Run the evaluation from scratch. Default is False.",
)
parser.add_argument(
"--semantic-description",
action=argparse.BooleanOptionalAction,
help="Whether to add semantic description before asking. Default is False.",
)
parser.add_argument(
"--custom-semantic-description",
action=argparse.BooleanOptionalAction,
help="Whether to add customized semantic description before asking. Default is False.",
)
parser.add_argument(
"--without-db-schema",
action=argparse.BooleanOptionalAction,
help="Whether to exclude the database schema information. Default is False.",
)
parser.add_argument(
"--easy-questions",
action=argparse.BooleanOptionalAction,
help="Whether to use easy questions for evaluation. Default is False.",
)
parser.add_argument(
"--hard-questions",
action=argparse.BooleanOptionalAction,
help="Whether to use hard questions for evaluation. Default is False.",
)
parser.add_argument
args = parser.parse_args()
PREDICTION_RESULTS_FILE = args.input_file
EVAL_AFTER_PREDICTION = args.eval_after_prediction
EVAL_FROM_SCRATCH = args.eval_from_scratch
ENABLE_SEMANTIC_DESCRIPTION = args.semantic_description
CUSTOM_SEMANTIC_DESCRIPTION = args.custom_semantic_description
NO_DB_SCHEMA = args.without_db_schema
EASY_QUESTIONS = args.easy_questions
HARD_QUESTIONS = args.hard_questions
with open(f"./src/eval/data/{DATASET_NAME}_data.json", "r") as f:
ground_truths = [json.loads(line) for line in f]
assert not (
CUSTOM_SEMANTIC_DESCRIPTION and ENABLE_SEMANTIC_DESCRIPTION
), "Cannot use both custom and general semantic description for evaluation."
assert not (
EASY_QUESTIONS and HARD_QUESTIONS
), "Cannot use both easy and hard questions for evaluation."
if EASY_QUESTIONS:
with open(f"./src/eval/data/{DATASET_NAME}_data_easy.json", "r") as f:
ground_truths = [json.loads(line) for line in f]
elif HARD_QUESTIONS:
with open(f"./src/eval/data/{DATASET_NAME}_data_hard.json", "r") as f:
ground_truths = [json.loads(line) for line in f]
else:
with open(f"./src/eval/data/{DATASET_NAME}_data.json", "r") as f:
ground_truths = [json.loads(line) for line in f]
if ENABLE_SEMANTIC_DESCRIPTION:
if os.path.exists(f"./src/eval/data/{DATASET_NAME}_with_semantic_mdl.json"):
print(f"Use the existed {DATASET_NAME}_with_semantic_mdl.json...\n")
else:
print(
f"Generating semantic description for the {DATASET_NAME} dataset...\n"
)
semantics_service = SemanticsService(
pipelines={
"generate_description": description.Generation(),
}
)
with open(f"./src/eval/data/{DATASET_NAME}_mdl.json", "r") as f:
mdl_data = json.load(f)
for model in tqdm(mdl_data["models"]):
semantic_desc = semantics_service.generate_description(
GenerateDescriptionRequest(
mdl=model,
model=model["name"],
identifier="model",
)
)
model["properties"]["description"] = semantic_desc.description
model["properties"]["display_name"] = semantic_desc.display_name
for column in model["columns"]:
semantic_desc = semantics_service.generate_description(
GenerateDescriptionRequest(
mdl=model,
model=model["name"],
identifier="column@" + column["name"],
)
)
column["properties"]["description"] = semantic_desc.description
column["properties"]["display_name"] = semantic_desc.display_name
with open(
f"./src/eval/data/{DATASET_NAME}_with_semantic_mdl.json", "w"
) as f:
json.dump(mdl_data, f)
if not Path("./outputs/ask").exists():
Path("./outputs/ask").mkdir(parents=True)
print(f"Running ask pipeline evaluation for the {DATASET_NAME} dataset...\n")
if (
PREDICTION_RESULTS_FILE
and Path(PREDICTION_RESULTS_FILE).exists()
and not args.eval_from_scratch
and not EVAL_FROM_SCRATCH
):
eval(Path(PREDICTION_RESULTS_FILE), DATASET_NAME, ground_truths)
else:
with open(f"./src/eval/data/{DATASET_NAME}_mdl.json", "r") as f:
mdl_str = json.dumps(json.load(f))
if ENABLE_SEMANTIC_DESCRIPTION:
with open(
f"./src/eval/data/{DATASET_NAME}_with_semantic_mdl.json", "r"
) as f:
mdl_str = json.dumps(json.load(f))
elif CUSTOM_SEMANTIC_DESCRIPTION:
with open(
f"./src/eval/data/{DATASET_NAME}_custom_semantic_mdl.json", "r"
) as f:
mdl_str = json.dumps(json.load(f))
else:
with open(f"./src/eval/data/{DATASET_NAME}_mdl.json", "r") as f:
mdl_str = json.dumps(json.load(f))
document_store = init_document_store(
dataset_name=DATASET_NAME,
recreate_index=True,
)
embedder = init_embedder(with_trace=with_trace)
embedder = init_embedder()
retriever = init_retriever(
document_store=document_store,
with_trace=with_trace,
top_k=10,
)
text_to_sql_generator = init_generator(
with_trace=with_trace,
)
sql_correction_generator = init_generator(
with_trace=with_trace,
)
text_to_sql_generator = init_generator()
sql_correction_generator = init_generator()
print("Indexing documents...")
indexing_pipeline = Indexing(document_store=document_store)
@@ -266,29 +363,25 @@ if __name__ == "__main__":
retrieval_pipeline = Retrieval(
embedder=embedder,
retriever=retriever,
with_trace=with_trace,
)
retrieval_pipeline_def = retrieval_pipeline._pipe.dumps()
generation_pipeline = Generation(
text_to_sql_generator=text_to_sql_generator,
with_trace=with_trace,
generator=text_to_sql_generator,
)
generation_pipeline_def = generation_pipeline._pipe.dumps()
sql_correction_pipeline = SQLCorrection(
sql_correction_generator=sql_correction_generator,
generator=sql_correction_generator,
)
sql_correction_pipeline_def = sql_correction_pipeline._pipe.dumps()
print(f"Running predictions for {len(ground_truths)} questions...")
start = time.time()
user_id = str(uuid.uuid4())
max_workers = os.cpu_count() // 2 if with_trace else None
user_id = str(uuid.uuid4()) if with_trace else None
with ThreadPoolExecutor(max_workers=max_workers) as executor:
with ThreadPoolExecutor() as executor:
args_list = [
(ground_truth["question"], user_id) for ground_truth in ground_truths
(ground_truth["question"], NO_DB_SCHEMA)
for ground_truth in ground_truths
]
outputs = list(
tqdm(
@@ -300,8 +393,11 @@ if __name__ == "__main__":
print(f"Time taken: {end - start:.2f}s")
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
print(
f"Write predictions to ./outputs/ask/{DATASET_NAME}_predictions_{timestamp}.json"
)
write_prediction_results(
f"./outputs/{DATASET_NAME}_predictions_{timestamp}.json",
f"./outputs/ask/{DATASET_NAME}_predictions_{timestamp}.json",
ground_truths,
outputs,
{
@@ -311,12 +407,10 @@ if __name__ == "__main__":
"sql_correction": sql_correction_pipeline_def,
},
)
if with_trace:
langfuse.flush()
if EVAL_AFTER_PREDICTION:
eval(
Path(f"./outputs/{DATASET_NAME}_predictions_{timestamp}.json"),
Path(f"./outputs/ask/{DATASET_NAME}_predictions_{timestamp}.json"),
DATASET_NAME,
ground_truths,
)
+815
View File
@@ -0,0 +1,815 @@
import json
import os
import re
import sqlite3
import subprocess
import zipfile
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List
import gdown
import pandas as pd
import sqlglot
import sqlparse
from dotenv import load_dotenv
from tqdm import tqdm
from tqdm.contrib import tzip
from src.eval.utils import get_generation_model_pricing
load_dotenv(override=True)
def semantic_diff(sql_query1: str, sql_query2: str):
try:
diff = sqlglot.diff(
sqlglot.parse_one(sql_query1, read=sqlglot.Dialects.TRINO),
sqlglot.parse_one(sql_query2, read=sqlglot.Dialects.TRINO),
)
for d in diff:
if str(d).startswith("Keep"):
continue
return True
return False
except Exception as e:
print(f"semantic_diff: {e}")
return True
def execute_sql_query(sql_query: str, db_path: str):
conn = sqlite3.connect(db_path)
cur = conn.cursor()
try:
cur.execute(sql_query)
# make each row a tuple of strings for easier comparison with the results from wren-engine
# also sort each row to make the order of the columns consistent
results = tuple(tuple(sorted(map(str, row))) for row in cur.fetchall())
except Exception:
results = []
finally:
cur.close()
conn.close()
return results
def execute_sql_query_through_wren_engine(sql_query: str):
command = f'psql -d "postgres://localhost:7432/canner-cml?options=--search_path%3Dspider" -c "{sql_query}"'
process = subprocess.Popen(
command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True
)
output, error = process.communicate()
if error:
return "", error.decode()
df = get_csv_table_from_response(output.decode())
# also sort each row to make the order of the columns consistent
sorted_df = df.apply(lambda x: tuple(sorted(x)), axis=1)
return tuple(sorted_df.values.tolist()), ""
def get_csv_table_from_response(full_response: str):
lines = [line.strip() for line in full_response.strip().split("\n")[2:-1]]
table_content = [[element.strip() for element in line.split("|")] for line in lines]
sql_query_result_in_csv = pd.DataFrame(table_content)
return sql_query_result_in_csv
def ground_truth_query_results_issubset(
ground_truth_query_results, prediction_query_results
):
rows1 = sorted(frozenset(row) for row in ground_truth_query_results)
rows2 = sorted(frozenset(row) for row in prediction_query_results)
for row1, row2 in zip(rows1, rows2):
if not row1.issubset(row2):
return False
return True
def get_ragas_eval_results(
ragas_eval_results: Dict[str, Any],
index: int,
):
return {
metric: eval_result["results"][index][0]["score"]
for metric, eval_result in ragas_eval_results.items()
}
# TODO: refactor this function if possible in the future
def generate_eval_report(
database_name: str,
groundtruths: List[Dict[str, str]],
predictions: List[Dict[str, Any]],
ragas_eval_results: Dict[str, Any],
):
results = {
"eval_results": {
"average_accuracy": 0,
"average_generation_cost": {
"total": 0,
"text_to_sql": {
"input": 0,
"output": 0,
},
"sql_correction": {
"input": 0,
"output": 0,
},
},
"average_latency": {
"total": 0,
"retrieval": 0,
"generation": {
"text_to_sql": 0,
"sql_correction": 0,
},
},
"sql_statistics": {
"text_to_sql": {
"valid": 0,
"invalid": 0,
"empty": 0,
},
"sql_correction": {
"valid": 0,
"invalid": 0,
},
},
"number_of_failed_queries": 0,
"details": {
"correct": {
"sql_semantic_same": [],
"query_results_same": [],
"ground_truth_query_results_issubset": [],
},
"wrong": [],
},
},
"pipelines": [],
}
if len(predictions) > 0:
results["pipelines"] = predictions[0]["pipelines"]
total = 0
correct = 0
for i, (ground_truth, prediction) in enumerate(tzip(groundtruths, predictions)):
total += 1
## dealing with cost part
for pipeline, pipeline_metadata in prediction["metadata"]["generation"].items():
if pipeline_metadata:
model_name = pipeline_metadata["model"]
generation_model_pricing = get_generation_model_pricing(model_name)
results["eval_results"]["average_generation_cost"][pipeline][
"input"
] += (
generation_model_pricing["prompt_tokens"]
* pipeline_metadata["usage"]["prompt_tokens"]
)
results["eval_results"]["average_generation_cost"][pipeline][
"output"
] += (
generation_model_pricing["completion_tokens"]
* pipeline_metadata["usage"]["completion_tokens"]
)
## dealing with latency part
results["eval_results"]["average_latency"]["retrieval"] += prediction[
"metadata"
]["latency"]["retrieval"]
for pipeline, pipeline_metadata in prediction["metadata"]["latency"][
"generation"
].items():
results["eval_results"]["average_latency"]["generation"][
pipeline
] += pipeline_metadata
## dealing with sql statistics part
for pipeline, values in prediction["sql_statistics"].items():
for key, value in values.items():
results["eval_results"]["sql_statistics"][pipeline][key] += value
## dealing with accuracy part
assert ground_truth["question"] == prediction["question"]
question = ground_truth["question"]
ground_truth_query_results = []
prediction_query_results = []
prediction_error_details = []
# directly compare the sql query using semantic diff
if prediction["answer"]:
if not semantic_diff(ground_truth["answer"], prediction["answer"]):
correct += 1
results["eval_results"]["details"]["correct"][
"sql_semantic_same"
].append(
{
"question": ground_truth["question"],
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
continue
# since the order of the sql query may be different, we compare the results as sets
ground_truth_query_results = execute_sql_query(
ground_truth["answer"],
f"./spider/database/{database_name}/{database_name}.sqlite",
)
(
prediction_query_results,
prediction_error_details,
) = execute_sql_query_through_wren_engine(
prediction["answer"],
)
if prediction_error_details:
results["eval_results"]["number_of_failed_queries"] += 1
if len(ground_truth_query_results) == len(prediction_query_results):
if set(ground_truth_query_results) == set(prediction_query_results):
correct += 1
results["eval_results"]["details"]["correct"][
"query_results_same"
].append(
{
"question": question,
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
continue
elif ground_truth_query_results_issubset(
ground_truth_query_results, prediction_query_results
):
correct += 1
results["eval_results"]["details"]["correct"][
"ground_truth_query_results_issubset"
].append(
{
"question": question,
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ground_truth_query_results": ground_truth_query_results,
"prediction_query_results": prediction_query_results,
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
continue
else:
results["eval_results"]["number_of_failed_queries"] += 1
results["eval_results"]["details"]["wrong"].append(
{
"question": question,
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ground_truth_query_results": ground_truth_query_results,
"prediction_query_results": prediction_query_results,
"prediction_error_details": prediction_error_details,
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
results["eval_results"]["average_accuracy"] = correct / total
results["eval_results"]["average_generation_cost"]["total"] = (
results["eval_results"]["average_generation_cost"]["text_to_sql"]["input"]
+ results["eval_results"]["average_generation_cost"]["text_to_sql"]["output"]
+ results["eval_results"]["average_generation_cost"]["sql_correction"]["input"]
+ results["eval_results"]["average_generation_cost"]["sql_correction"]["output"]
) / total
results["eval_results"]["average_generation_cost"]["text_to_sql"]["input"] /= total
results["eval_results"]["average_generation_cost"]["text_to_sql"]["output"] /= total
results["eval_results"]["average_generation_cost"]["sql_correction"][
"input"
] /= total
results["eval_results"]["average_generation_cost"]["sql_correction"][
"output"
] /= total
results["eval_results"]["average_latency"]["total"] = (
results["eval_results"]["average_latency"]["retrieval"]
+ results["eval_results"]["average_latency"]["generation"]["text_to_sql"]
+ results["eval_results"]["average_latency"]["generation"]["sql_correction"]
) / total
results["eval_results"]["average_latency"]["retrieval"] /= total
results["eval_results"]["average_latency"]["generation"]["text_to_sql"] /= total
results["eval_results"]["average_latency"]["generation"]["sql_correction"] /= total
return results
def download_spider_data():
if Path("spider").exists():
return
if Path("spider.zip").exists():
os.remove("spider.zip")
# 1. uploaded to Jimmy's google drive from the official Spider dataset website(data based at 2024/01/28)
# 2. added `table_counts_in_database.json`
# 3. changed column `salary` to `player_salary` in `baseball_1.player` table,
# since in bigquery, column name should not be the same as table name
url = "https://drive.google.com/u/0/uc?id=1StzD_Yha1W-BJOLimuvdzH-cF6sEc_ak&export=download"
output = "spider.zip"
gdown.download(url, output, quiet=False)
with zipfile.ZipFile(output, "r") as zip_ref:
zip_ref.extractall(".")
os.remove("spider.zip")
def write_prediction_results(
file_name: str,
ground_truths: List[Dict],
outputs: List[Dict[str, Any]],
pipeline_defs: Dict[str, str],
):
with open(file_name, "w") as f:
for ground_truth, output in tzip(ground_truths, outputs):
json.dump(
{
"question": ground_truth["question"],
"contexts": [
{
"content": [context.content],
"score": context.score,
}
for context in output["contexts"]
],
"answer": output["prediction"],
"metadata": output["metadata"],
"sql_statistics": output["sql_statistics"],
"pipelines": pipeline_defs,
},
f,
)
f.write("\n")
def prepare_evaluation_pipeline_inputs(
component_names: List[str],
ground_truths_data: List[Dict[str, Any]],
predictions_data: List[Dict[str, Any]],
):
inputs = {}
questions = []
ground_truths = []
contexts = []
responses = []
for ground_truth, prediction in tzip(ground_truths_data, predictions_data):
assert ground_truth["question"] == prediction["question"]
questions.append(ground_truth["question"])
ground_truths.append(ground_truth["answer"])
contexts.append(
[json.dumps(context["content"]) for context in prediction["contexts"]]
)
responses.append(prediction["answer"])
for component_name in tqdm(component_names):
ragas_metric = "_".join(component_name.split("_")[1:])
# https://docs.haystack.deepset.ai/v2.0/docs/ragasevaluator#supported-metrics
if ragas_metric == "ANSWER_CORRECTNESS":
inputs[component_name] = {
"questions": questions,
"responses": responses,
"ground_truths": ground_truths,
}
elif ragas_metric == "FAITHFULNESS":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
elif ragas_metric == "ANSWER_SIMILARITY":
inputs[component_name] = {
"responses": responses,
"ground_truths": ground_truths,
}
elif ragas_metric == "CONTEXT_PRECISION":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"ground_truths": ground_truths,
}
elif ragas_metric == "CONTEXT_UTILIZATION":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
elif ragas_metric == "CONTEXT_RECALL":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"ground_truths": ground_truths,
}
elif ragas_metric == "ASPECT_CRITIQUE":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
elif ragas_metric == "CONTEXT_RELEVANCY":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
}
elif ragas_metric == "ANSWER_RELEVANCY":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
return inputs
def get_latest_prediction_outputs_file(path: Path, dataset_name: str) -> str:
def _extract_datetime(file_name: Path) -> datetime:
file_name, _ = os.path.splitext(file_name)
timestamp_str = file_name.split("_")[-1]
return datetime.strptime(timestamp_str, "%Y%m%d%H%M%S")
files = list(path.glob(f"{dataset_name}_predictions*.json"))
if not files:
return ""
return str(sorted(files, key=_extract_datetime, reverse=True)[0])
def transpile_sql_from_sqlite_to_trino(sql_query: str):
return sqlglot.transpile(
sql_query, read=sqlglot.Dialects.SQLITE, write=sqlglot.Dialects.TRINO
)[0]
def get_database_names() -> list[str]:
return [
folder.name for folder in Path("spider/database").iterdir() if folder.is_dir()
]
def get_table_names(db_path: str) -> list[str]:
# Connect to the SQLite database
conn = sqlite3.connect(db_path)
# Create a cursor object
cur = conn.cursor()
# Get the table names
cur.execute("SELECT name FROM sqlite_master WHERE type='table' ORDER BY name")
# Fetch the results
results = cur.fetchall()
cur.close()
conn.close()
return [result[0] for result in results]
def get_database_schema(
db_path: str, table_names: list[str], should_save_file: bool = False
) -> list[dict]:
# Connect to the SQLite database
conn = sqlite3.connect(db_path)
# Create a cursor object
cur = conn.cursor()
# Get the table schemas
table_schemas = []
for table_name in table_names:
cur.execute(
f"SELECT sql FROM sqlite_master WHERE type='table' AND name='{table_name}'"
)
row = cur.fetchone()
if row is not None:
table_schema = row[0]
table_schemas.append(
{
"table_name": table_name,
"table_schema": re.sub(r"\s+", " ", table_schema).strip(),
}
)
else:
print(f"No table named {table_name} found.")
cur.close()
if should_save_file:
file_name = db_path.split("/")[-1].split(".")[0]
with open(f"{file_name}_schema.txt", "w") as file:
for table_schema in table_schemas:
file.write(f"Table name: {table_schema['table_name']}\n")
file.write(f"Table schema: {table_schema['table_schema']}\n\n")
return table_schemas
def get_table_relationships(db_path: str):
# Connect to the SQLite database
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
# Get a list of tables in the database
cursor.execute("SELECT name FROM sqlite_master WHERE type='table';")
tables = [row[0] for row in cursor.fetchall()]
# Function to check if a column is part of a unique or primary key
def is_unique_or_pk(table, column):
cursor.execute(f"PRAGMA table_info('{table}')")
for col in cursor.fetchall():
# Check if the column is a primary key or part of a unique constraint
if col[1] == column and (col[5] > 0 or col[3]):
return True
return False
# Analyze relationships
relationships = {}
for table in tables:
cursor.execute(f"PRAGMA foreign_key_list('{table}')")
fk_list = cursor.fetchall()
if not fk_list:
continue
for fk in fk_list:
ref_table = fk[2]
from_column = fk[3]
to_column = fk[4]
# Determine relationship type
if is_unique_or_pk(table, from_column):
if is_unique_or_pk(ref_table, to_column):
relation_type = "ONE_TO_ONE"
else:
relation_type = "ONE_TO_MANY"
else:
if is_unique_or_pk(ref_table, to_column):
relation_type = "MANY_TO_ONE"
else:
relation_type = "MANY_TO_MANY"
relationships[(table, ref_table)] = relation_type
conn.close()
return relationships
def get_appropriat_column_type(column_type: str):
if column_type.lower() == "text" or "varchar" in column_type.lower():
return "VARCHAR"
elif column_type.lower() == "numeric":
return "REAL"
elif column_type.lower() == "int":
return "INTEGER"
return column_type.upper()
def split_table_definition(table_definition: str):
return table_definition.split(", ")
def parse_column_definition(column_definition: str):
column_def = column_definition.split(" ")
return {
"name": column_def[0],
"type": column_def[1] if len(column_def) > 1 else "TEXT",
"not_null": True
if len(column_def) == 3 and column_def[2].lower() == "not null"
else False,
}
def parse_table_definition(
table_name: str, table_definition: str, relationships_info: dict
):
match = re.search(r"\((.*)\)", table_definition)
assert match
inside_parentheses = match.group(1)
parts = split_table_definition(inside_parentheses)
# Lists to store columns and foreign keys
columns = []
relationships = []
primary_key = ""
for part in parts:
if part.startswith("foreign key") or part.startswith("FOREIGN KEY"):
part = part.replace("`", "").replace('"', "")
regex1 = r"FOREIGN KEY\(([^)]+)\) REFERENCES ([^(]+)\(([^)]+)\)"
regex2 = r"foreign key\(([^)]+)\) references ([^(]+)\(([^)]+)\)"
match = re.search(regex1, part)
match2 = re.search(regex2, part)
match_result = False
if match:
match_result = match
elif match2:
match_result = match2
if match_result:
relationships.append(
{
"table_name": table_name,
"foreign_key": match_result.group(1),
"ref_table_name": match_result.group(2),
"ref_column": match_result.group(3),
"join_type": relationships_info[
(table_name, match_result.group(2))
],
"properties": {},
}
)
else:
if "PRIMARY KEY" in part or "primary key" in part:
primary_key = part.strip().split(" ")[0]
part = (
part.replace("PRIMARY KEY", "").replace("primary key", "").strip()
)
# Splitting the column name and type
column_def = parse_column_definition(part.strip())
columns.append(
{
"name": column_def["name"].replace('"', ""),
"type": get_appropriat_column_type(column_def["type"]),
"notNull": column_def[
"not_null"
], # Assuming notNull is False by default as not specified in the string
"isCalculated": False, # Assuming isCalculated is False by default
"expression": column_def["name"].replace(
'"', ""
), # Assuming expression is the column name itself
"properties": {},
}
)
if relationships:
for relationship in relationships:
columns.append(
{
"name": relationship["ref_table_name"],
"type": relationship["ref_table_name"],
"notNull": True,
"isCalculated": False,
"relationship": f"{relationship['table_name']}_{relationship['ref_table_name']}",
"properties": {},
}
)
return {
"columns": columns,
"primary_key": primary_key,
"relationships": relationships,
}
def generate_mdl_json(
database_schema: list[dict],
catalog_name: str,
schema_name: str,
database_name: str,
relationships_info: dict,
should_save_file: bool = False,
file_path: str = "",
):
mdl_json = {
"catalog": catalog_name,
"schema": schema_name,
"models": [],
"relationships": [],
# these will be empty for now
"metrics": [],
"cumulativeMetrics": [],
"enumDefinitions": [],
"views": [],
"macros": [],
}
for table in database_schema:
# remove comments
clean_table_schema = table["table_schema"].replace(
"-- this should be removed", ""
)
parsed = sqlparse.parse(clean_table_schema)[0]
table_definition = parse_table_definition(
table["table_name"], str(parsed.tokens[-1]), relationships_info
)
mdl_json["models"].append(
{
"name": table["table_name"],
"properties": {},
"refSql": f"select * from \"{catalog_name}\".{schema_name}.\"{database_name}-{table['table_name']}\"",
"columns": table_definition["columns"],
"primaryKey": table_definition["primary_key"],
}
)
if table_definition["relationships"]:
for relationship in table_definition["relationships"]:
mdl_json["relationships"].append(
{
"name": f"{relationship['table_name']}_{relationship['ref_table_name']}",
"models": [
relationship["table_name"],
relationship["ref_table_name"],
],
"joinType": relationship["join_type"],
"condition": f"{relationship['table_name']}.{relationship['foreign_key']} = {relationship['ref_table_name']}.{relationship['ref_column']}",
}
)
if should_save_file:
data_root = "src/eval/data"
if not Path(data_root).exists():
Path(data_root).mkdir(parents=True, exist_ok=True)
if not file_path:
file_path = f"{data_root}/{database_name}.json"
# save the file
with open(file_path, "w") as file:
json.dump(mdl_json, file, indent=2)
print(
f"MDL JSON for {database_name} generated successfully. Check the {data_root} folder."
)
return mdl_json
def generate_text_to_sql_dataset(
paths: list[str],
database_name: str = "college_3",
should_save_file: bool = False,
file_path: str = "data/college_3_data.json",
):
target_data = []
for path in paths:
with open(path, "r") as f:
data = json.load(f)
if database_name == "":
target_data += data
else:
for entry in data:
if entry["db_id"] == database_name:
target_data.append(
{
"question": entry["question"],
"answer": transpile_sql_from_sqlite_to_trino(
re.sub(r"\s+", " ", entry["query"]).strip()
),
}
)
if should_save_file:
data_root = "src/eval/data"
if not Path(data_root).exists():
Path(data_root).mkdir(parents=True, exist_ok=True)
with open(file_path, "w") as f:
for entry in target_data:
json.dump(entry, f)
f.write("\n")
print(
f"Dataset for {database_name} is generated successfully. Check the {data_root} folder."
)
return target_data
-815
View File
@@ -1,107 +1,3 @@
import json
import os
import re
import sqlite3
import subprocess
import zipfile
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List
import gdown
import pandas as pd
import sqlglot
import sqlparse
from dotenv import load_dotenv
from tqdm import tqdm
from tqdm.contrib import tzip
load_dotenv(override=True)
def semantic_diff(sql_query1: str, sql_query2: str):
try:
diff = sqlglot.diff(
sqlglot.parse_one(sql_query1, read=sqlglot.Dialects.TRINO),
sqlglot.parse_one(sql_query2, read=sqlglot.Dialects.TRINO),
)
for d in diff:
if str(d).startswith("Keep"):
continue
return True
return False
except Exception as e:
print(f"semantic_diff: {e}")
return True
def execute_sql_query(sql_query: str, db_path: str):
conn = sqlite3.connect(db_path)
cur = conn.cursor()
try:
cur.execute(sql_query)
# make each row a tuple of strings for easier comparison with the results from wren-engine
# also sort each row to make the order of the columns consistent
results = tuple(tuple(sorted(map(str, row))) for row in cur.fetchall())
except Exception:
results = []
finally:
cur.close()
conn.close()
return results
def execute_sql_query_through_wren_engine(sql_query: str):
command = f'psql -d "postgres://localhost:7432/canner-cml?options=--search_path%3Dspider" -c "{sql_query}"'
process = subprocess.Popen(
command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, shell=True
)
output, error = process.communicate()
if error:
return "", error.decode()
df = get_csv_table_from_response(output.decode())
# also sort each row to make the order of the columns consistent
sorted_df = df.apply(lambda x: tuple(sorted(x)), axis=1)
return tuple(sorted_df.values.tolist()), ""
def get_csv_table_from_response(full_response: str):
lines = [line.strip() for line in full_response.strip().split("\n")[2:-1]]
table_content = [[element.strip() for element in line.split("|")] for line in lines]
sql_query_result_in_csv = pd.DataFrame(table_content)
return sql_query_result_in_csv
def ground_truth_query_results_issubset(
ground_truth_query_results, prediction_query_results
):
rows1 = sorted(frozenset(row) for row in ground_truth_query_results)
rows2 = sorted(frozenset(row) for row in prediction_query_results)
for row1, row2 in zip(rows1, rows2):
if not row1.issubset(row2):
return False
return True
def get_ragas_eval_results(
ragas_eval_results: Dict[str, Any],
index: int,
):
return {
metric: eval_result["results"][index][0]["score"]
for metric, eval_result in ragas_eval_results.items()
}
def get_generation_model_pricing(
model_name: str,
):
@@ -126,714 +22,3 @@ def get_generation_model_pricing(
}
return generation_model_pricing[model_name]
# TODO: refactor this function if possible in the future
def generate_eval_report(
database_name: str,
groundtruths: List[Dict[str, str]],
predictions: List[Dict[str, Any]],
ragas_eval_results: Dict[str, Any],
):
results = {
"eval_results": {
"average_accuracy": 0,
"average_generation_cost": {
"total": 0,
"text_to_sql": {
"input": 0,
"output": 0,
},
"sql_correction": {
"input": 0,
"output": 0,
},
},
"average_latency": {
"total": 0,
"retrieval": 0,
"generation": {
"text_to_sql": 0,
"sql_correction": 0,
},
},
"sql_statistics": {
"text_to_sql": {
"valid": 0,
"invalid": 0,
"empty": 0,
},
"sql_correction": {
"valid": 0,
"invalid": 0,
},
},
"number_of_failed_queries": 0,
"details": {
"correct": {
"sql_semantic_same": [],
"query_results_same": [],
"ground_truth_query_results_issubset": [],
},
"wrong": [],
},
},
"pipelines": [],
}
if len(predictions) > 0:
results["pipelines"] = predictions[0]["pipelines"]
total = 0
correct = 0
for i, (ground_truth, prediction) in enumerate(tzip(groundtruths, predictions)):
total += 1
## dealing with cost part
for pipeline, pipeline_metadata in prediction["metadata"]["generation"].items():
if pipeline_metadata:
model_name = pipeline_metadata["model"]
generation_model_pricing = get_generation_model_pricing(model_name)
results["eval_results"]["average_generation_cost"][pipeline][
"input"
] += (
generation_model_pricing["prompt_tokens"]
* pipeline_metadata["usage"]["prompt_tokens"]
)
results["eval_results"]["average_generation_cost"][pipeline][
"output"
] += (
generation_model_pricing["completion_tokens"]
* pipeline_metadata["usage"]["completion_tokens"]
)
## dealing with latency part
results["eval_results"]["average_latency"]["retrieval"] += prediction[
"metadata"
]["latency"]["retrieval"]
for pipeline, pipeline_metadata in prediction["metadata"]["latency"][
"generation"
].items():
results["eval_results"]["average_latency"]["generation"][
pipeline
] += pipeline_metadata
## dealing with sql statistics part
for pipeline, values in prediction["sql_statistics"].items():
for key, value in values.items():
results["eval_results"]["sql_statistics"][pipeline][key] += value
## dealing with accuracy part
assert ground_truth["question"] == prediction["question"]
question = ground_truth["question"]
ground_truth_query_results = []
prediction_query_results = []
prediction_error_details = []
# directly compare the sql query using semantic diff
if prediction["answer"]:
if not semantic_diff(ground_truth["answer"], prediction["answer"]):
correct += 1
results["eval_results"]["details"]["correct"][
"sql_semantic_same"
].append(
{
"question": ground_truth["question"],
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
continue
# since the order of the sql query may be different, we compare the results as sets
ground_truth_query_results = execute_sql_query(
ground_truth["answer"],
f"./spider/database/{database_name}/{database_name}.sqlite",
)
(
prediction_query_results,
prediction_error_details,
) = execute_sql_query_through_wren_engine(
prediction["answer"],
)
if prediction_error_details:
results["eval_results"]["number_of_failed_queries"] += 1
if len(ground_truth_query_results) == len(prediction_query_results):
if set(ground_truth_query_results) == set(prediction_query_results):
correct += 1
results["eval_results"]["details"]["correct"][
"query_results_same"
].append(
{
"question": question,
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
continue
elif ground_truth_query_results_issubset(
ground_truth_query_results, prediction_query_results
):
correct += 1
results["eval_results"]["details"]["correct"][
"ground_truth_query_results_issubset"
].append(
{
"question": question,
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ground_truth_query_results": ground_truth_query_results,
"prediction_query_results": prediction_query_results,
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
continue
else:
results["eval_results"]["number_of_failed_queries"] += 1
results["eval_results"]["details"]["wrong"].append(
{
"question": question,
"ground_truth_answer": ground_truth["answer"],
"prediction_answer": prediction["answer"],
"ground_truth_query_results": ground_truth_query_results,
"prediction_query_results": prediction_query_results,
"prediction_error_details": prediction_error_details,
"ragas_eval_results": get_ragas_eval_results(
ragas_eval_results,
i,
),
}
)
results["eval_results"]["average_accuracy"] = correct / total
results["eval_results"]["average_generation_cost"]["total"] = (
results["eval_results"]["average_generation_cost"]["text_to_sql"]["input"]
+ results["eval_results"]["average_generation_cost"]["text_to_sql"]["output"]
+ results["eval_results"]["average_generation_cost"]["sql_correction"]["input"]
+ results["eval_results"]["average_generation_cost"]["sql_correction"]["output"]
) / total
results["eval_results"]["average_generation_cost"]["text_to_sql"]["input"] /= total
results["eval_results"]["average_generation_cost"]["text_to_sql"]["output"] /= total
results["eval_results"]["average_generation_cost"]["sql_correction"][
"input"
] /= total
results["eval_results"]["average_generation_cost"]["sql_correction"][
"output"
] /= total
results["eval_results"]["average_latency"]["total"] = (
results["eval_results"]["average_latency"]["retrieval"]
+ results["eval_results"]["average_latency"]["generation"]["text_to_sql"]
+ results["eval_results"]["average_latency"]["generation"]["sql_correction"]
) / total
results["eval_results"]["average_latency"]["retrieval"] /= total
results["eval_results"]["average_latency"]["generation"]["text_to_sql"] /= total
results["eval_results"]["average_latency"]["generation"]["sql_correction"] /= total
return results
def download_spider_data():
if Path("spider").exists():
return
if Path("spider.zip").exists():
os.remove("spider.zip")
# 1. uploaded to Jimmy's google drive from the official Spider dataset website(data based at 2024/01/28)
# 2. added `table_counts_in_database.json`
# 3. changed column `salary` to `player_salary` in `baseball_1.player` table,
# since in bigquery, column name should not be the same as table name
url = "https://drive.google.com/u/0/uc?id=1StzD_Yha1W-BJOLimuvdzH-cF6sEc_ak&export=download"
output = "spider.zip"
gdown.download(url, output, quiet=False)
with zipfile.ZipFile(output, "r") as zip_ref:
zip_ref.extractall(".")
os.remove("spider.zip")
def write_prediction_results(
file_name: str,
ground_truths: List[Dict],
outputs: List[Dict[str, Any]],
pipeline_defs: Dict[str, str],
):
with open(file_name, "w") as f:
for ground_truth, output in tzip(ground_truths, outputs):
json.dump(
{
"question": ground_truth["question"],
"contexts": [
{
"content": [json.loads(context.content)],
"score": context.score,
}
for context in output["contexts"]
],
"answer": output["prediction"],
"metadata": output["metadata"],
"sql_statistics": output["sql_statistics"],
"pipelines": pipeline_defs,
},
f,
)
f.write("\n")
def prepare_evaluation_pipeline_inputs(
component_names: List[str],
ground_truths_data: List[Dict[str, Any]],
predictions_data: List[Dict[str, Any]],
):
inputs = {}
questions = []
ground_truths = []
contexts = []
responses = []
for ground_truth, prediction in tzip(ground_truths_data, predictions_data):
assert ground_truth["question"] == prediction["question"]
questions.append(ground_truth["question"])
ground_truths.append(ground_truth["answer"])
contexts.append(
[json.dumps(context["content"]) for context in prediction["contexts"]]
)
responses.append(prediction["answer"])
for component_name in tqdm(component_names):
ragas_metric = "_".join(component_name.split("_")[1:])
# https://docs.haystack.deepset.ai/v2.0/docs/ragasevaluator#supported-metrics
if ragas_metric == "ANSWER_CORRECTNESS":
inputs[component_name] = {
"questions": questions,
"responses": responses,
"ground_truths": ground_truths,
}
elif ragas_metric == "FAITHFULNESS":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
elif ragas_metric == "ANSWER_SIMILARITY":
inputs[component_name] = {
"responses": responses,
"ground_truths": ground_truths,
}
elif ragas_metric == "CONTEXT_PRECISION":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"ground_truths": ground_truths,
}
elif ragas_metric == "CONTEXT_UTILIZATION":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
elif ragas_metric == "CONTEXT_RECALL":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"ground_truths": ground_truths,
}
elif ragas_metric == "ASPECT_CRITIQUE":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
elif ragas_metric == "CONTEXT_RELEVANCY":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
}
elif ragas_metric == "ANSWER_RELEVANCY":
inputs[component_name] = {
"questions": questions,
"contexts": contexts,
"responses": responses,
}
return inputs
def get_latest_prediction_outputs_file(path: Path, dataset_name: str) -> str:
def _extract_datetime(file_name: Path) -> datetime:
file_name, _ = os.path.splitext(file_name)
timestamp_str = file_name.split("_")[-1]
return datetime.strptime(timestamp_str, "%Y%m%d%H%M%S")
files = list(path.glob(f"{dataset_name}_predictions*.json"))
if not files:
return ""
return str(sorted(files, key=_extract_datetime, reverse=True)[0])
def transpile_sql_from_sqlite_to_trino(sql_query: str):
return sqlglot.transpile(
sql_query, read=sqlglot.Dialects.SQLITE, write=sqlglot.Dialects.TRINO
)[0]
def get_database_names() -> list[str]:
return [
folder.name for folder in Path("spider/database").iterdir() if folder.is_dir()
]
def get_table_names(db_path: str) -> list[str]:
# Connect to the SQLite database
conn = sqlite3.connect(db_path)
# Create a cursor object
cur = conn.cursor()
# Get the table names
cur.execute("SELECT name FROM sqlite_master WHERE type='table' ORDER BY name")
# Fetch the results
results = cur.fetchall()
cur.close()
conn.close()
return [result[0] for result in results]
def get_database_schema(
db_path: str, table_names: list[str], should_save_file: bool = False
) -> list[dict]:
# Connect to the SQLite database
conn = sqlite3.connect(db_path)
# Create a cursor object
cur = conn.cursor()
# Get the table schemas
table_schemas = []
for table_name in table_names:
cur.execute(
f"SELECT sql FROM sqlite_master WHERE type='table' AND name='{table_name}'"
)
row = cur.fetchone()
if row is not None:
table_schema = row[0]
table_schemas.append(
{
"table_name": table_name,
"table_schema": re.sub(r"\s+", " ", table_schema).strip(),
}
)
else:
print(f"No table named {table_name} found.")
cur.close()
if should_save_file:
file_name = db_path.split("/")[-1].split(".")[0]
with open(f"{file_name}_schema.txt", "w") as file:
for table_schema in table_schemas:
file.write(f"Table name: {table_schema['table_name']}\n")
file.write(f"Table schema: {table_schema['table_schema']}\n\n")
return table_schemas
def get_table_relationships(db_path: str):
# Connect to the SQLite database
conn = sqlite3.connect(db_path)
cursor = conn.cursor()
# Get a list of tables in the database
cursor.execute("SELECT name FROM sqlite_master WHERE type='table';")
tables = [row[0] for row in cursor.fetchall()]
# Function to check if a column is part of a unique or primary key
def is_unique_or_pk(table, column):
cursor.execute(f"PRAGMA table_info('{table}')")
for col in cursor.fetchall():
# Check if the column is a primary key or part of a unique constraint
if col[1] == column and (col[5] > 0 or col[3]):
return True
return False
# Analyze relationships
relationships = {}
for table in tables:
cursor.execute(f"PRAGMA foreign_key_list('{table}')")
fk_list = cursor.fetchall()
if not fk_list:
continue
for fk in fk_list:
ref_table = fk[2]
from_column = fk[3]
to_column = fk[4]
# Determine relationship type
if is_unique_or_pk(table, from_column):
if is_unique_or_pk(ref_table, to_column):
relation_type = "ONE_TO_ONE"
else:
relation_type = "ONE_TO_MANY"
else:
if is_unique_or_pk(ref_table, to_column):
relation_type = "MANY_TO_ONE"
else:
relation_type = "MANY_TO_MANY"
relationships[(table, ref_table)] = relation_type
conn.close()
return relationships
def get_appropriat_column_type(column_type: str):
if column_type.lower() == "text" or "varchar" in column_type.lower():
return "VARCHAR"
elif column_type.lower() == "numeric":
return "REAL"
elif column_type.lower() == "int":
return "INTEGER"
return column_type.upper()
def split_table_definition(table_definition: str):
return table_definition.split(", ")
def parse_column_definition(column_definition: str):
column_def = column_definition.split(" ")
return {
"name": column_def[0],
"type": column_def[1] if len(column_def) > 1 else "TEXT",
"not_null": True
if len(column_def) == 3 and column_def[2].lower() == "not null"
else False,
}
def parse_table_definition(
table_name: str, table_definition: str, relationships_info: dict
):
match = re.search(r"\((.*)\)", table_definition)
assert match
inside_parentheses = match.group(1)
parts = split_table_definition(inside_parentheses)
# Lists to store columns and foreign keys
columns = []
relationships = []
primary_key = ""
for part in parts:
if part.startswith("foreign key") or part.startswith("FOREIGN KEY"):
part = part.replace("`", "").replace('"', "")
regex1 = r"FOREIGN KEY\(([^)]+)\) REFERENCES ([^(]+)\(([^)]+)\)"
regex2 = r"foreign key\(([^)]+)\) references ([^(]+)\(([^)]+)\)"
match = re.search(regex1, part)
match2 = re.search(regex2, part)
match_result = False
if match:
match_result = match
elif match2:
match_result = match2
if match_result:
relationships.append(
{
"table_name": table_name,
"foreign_key": match_result.group(1),
"ref_table_name": match_result.group(2),
"ref_column": match_result.group(3),
"join_type": relationships_info[
(table_name, match_result.group(2))
],
"properties": {},
}
)
else:
if "PRIMARY KEY" in part or "primary key" in part:
primary_key = part.strip().split(" ")[0]
part = (
part.replace("PRIMARY KEY", "").replace("primary key", "").strip()
)
# Splitting the column name and type
column_def = parse_column_definition(part.strip())
columns.append(
{
"name": column_def["name"].replace('"', ""),
"type": get_appropriat_column_type(column_def["type"]),
"notNull": column_def[
"not_null"
], # Assuming notNull is False by default as not specified in the string
"isCalculated": False, # Assuming isCalculated is False by default
"expression": column_def["name"].replace(
'"', ""
), # Assuming expression is the column name itself
"properties": {},
}
)
if relationships:
for relationship in relationships:
columns.append(
{
"name": relationship["ref_table_name"],
"type": relationship["ref_table_name"],
"notNull": True,
"isCalculated": False,
"relationship": f"{relationship['table_name']}_{relationship['ref_table_name']}",
"properties": {},
}
)
return {
"columns": columns,
"primary_key": primary_key,
"relationships": relationships,
}
def generate_mdl_json(
database_schema: list[dict],
catalog_name: str,
schema_name: str,
database_name: str,
relationships_info: dict,
should_save_file: bool = False,
file_path: str = "",
):
mdl_json = {
"catalog": catalog_name,
"schema": schema_name,
"models": [],
"relationships": [],
# these will be empty for now
"metrics": [],
"cumulativeMetrics": [],
"enumDefinitions": [],
"views": [],
"macros": [],
}
for table in database_schema:
# remove comments
clean_table_schema = table["table_schema"].replace(
"-- this should be removed", ""
)
parsed = sqlparse.parse(clean_table_schema)[0]
table_definition = parse_table_definition(
table["table_name"], str(parsed.tokens[-1]), relationships_info
)
mdl_json["models"].append(
{
"name": table["table_name"],
"properties": {},
"refSql": f"select * from \"{catalog_name}\".{schema_name}.\"{database_name}-{table['table_name']}\"",
"columns": table_definition["columns"],
"primaryKey": table_definition["primary_key"],
}
)
if table_definition["relationships"]:
for relationship in table_definition["relationships"]:
mdl_json["relationships"].append(
{
"name": f"{relationship['table_name']}_{relationship['ref_table_name']}",
"models": [
relationship["table_name"],
relationship["ref_table_name"],
],
"joinType": relationship["join_type"],
"condition": f"{relationship['table_name']}.{relationship['foreign_key']} = {relationship['ref_table_name']}.{relationship['ref_column']}",
}
)
if should_save_file:
data_root = "src/eval/data"
if not Path(data_root).exists():
Path(data_root).mkdir(parents=True, exist_ok=True)
if not file_path:
file_path = f"{data_root}/{database_name}.json"
# save the file
with open(file_path, "w") as file:
json.dump(mdl_json, file, indent=2)
print(
f"MDL JSON for {database_name} generated successfully. Check the {data_root} folder."
)
return mdl_json
def generate_text_to_sql_dataset(
paths: list[str],
database_name: str = "college_3",
should_save_file: bool = False,
file_path: str = "data/college_3_data.json",
):
target_data = []
for path in paths:
with open(path, "r") as f:
data = json.load(f)
if database_name == "":
target_data += data
else:
for entry in data:
if entry["db_id"] == database_name:
target_data.append(
{
"question": entry["question"],
"answer": transpile_sql_from_sqlite_to_trino(
re.sub(r"\s+", " ", entry["query"]).strip()
),
}
)
if should_save_file:
data_root = "src/eval/data"
if not Path(data_root).exists():
Path(data_root).mkdir(parents=True, exist_ok=True)
with open(file_path, "w") as f:
for entry in target_data:
json.dump(entry, f)
f.write("\n")
print(
f"Dataset for {database_name} is generated successfully. Check the {data_root} folder."
)
return target_data
@@ -6,7 +6,7 @@ networks:
services:
wren-engine:
image: ghcr.io/canner/wren-engine:latest
image: ghcr.io/canner/wren-engine:nightly
pull_policy: always
platform: ${PLATFORM}
ports:
+13 -14
View File
@@ -1,7 +1,8 @@
import os
from dotenv import load_dotenv
from src.pipelines.ask import (
followup_generation_pipeline as ask_followup_generation_pipeline,
)
from src.pipelines.ask import (
generation_pipeline as ask_generation_pipeline,
)
@@ -39,17 +40,15 @@ ASK_DETAILS_SERVICE = None
def init_globals():
global SEMANTIC_SERVICE, ASK_SERVICE, ASK_DETAILS_SERVICE
with_trace = os.getenv("ENABLE_TRACE", default=False)
document_store = init_document_store()
embedder = init_embedder(with_trace=with_trace)
embedder = init_embedder()
retriever = init_retriever(
document_store=document_store,
with_trace=with_trace,
)
text_to_sql_generator = init_generator(with_trace=with_trace)
sql_correction_generator = init_generator(with_trace=with_trace)
sql_details_generator = init_ask_details_generator(with_trace=with_trace)
text_to_sql_generator = init_generator()
text_to_sql_with_followup_generator = init_generator()
sql_correction_generator = init_generator()
sql_details_generator = init_ask_details_generator()
SEMANTIC_SERVICE = SemanticsService(
pipelines={
@@ -64,14 +63,15 @@ def init_globals():
"retrieval": ask_retrieval_pipeline.Retrieval(
embedder=embedder,
retriever=retriever,
with_trace=with_trace,
),
"generation": ask_generation_pipeline.Generation(
text_to_sql_generator=text_to_sql_generator,
with_trace=with_trace,
generator=text_to_sql_generator,
),
"sql_correction": ask_sql_correction_pipeline.SQLCorrection(
sql_correction_generator=sql_correction_generator,
generator=sql_correction_generator,
),
"followup_generation": ask_followup_generation_pipeline.FollowUpGeneration(
generator=text_to_sql_with_followup_generator,
),
}
)
@@ -79,7 +79,6 @@ def init_globals():
pipelines={
"generation": ask_details_generation_pipeline.Generation(
generator=sql_details_generator,
with_trace=with_trace,
),
}
)
@@ -1,38 +1,15 @@
from typing import Any, Dict, List
from haystack import component
from haystack.components.embedders import OpenAITextEmbedder
from haystack.utils.auth import Secret
from src.utils import load_env_vars
from ...trace import TraceSpanInput, trace_span
load_env_vars()
EMBEDDING_MODEL_NAME = "text-embedding-3-large"
EMBEDDING_MODEL_DIMENSION = 3072
@component
class TracedOpenAITextEmbedder(OpenAITextEmbedder):
def _run(self, *args, **kwargs):
return super(TracedOpenAITextEmbedder, self).run(*args, **kwargs)
@component.output_types(embedding=List[float], meta=Dict[str, Any])
def run(self, text: str, trace_span_input: TraceSpanInput):
return trace_span(self._run)(trace_span_input=trace_span_input, text=text)
def init_embedder(
with_trace: bool = False, embedding_model_name: str = EMBEDDING_MODEL_NAME
):
if with_trace:
return TracedOpenAITextEmbedder(
api_key=Secret.from_env_var("OPENAI_API_KEY"),
model=embedding_model_name,
)
def init_embedder(embedding_model_name: str = EMBEDDING_MODEL_NAME):
return OpenAITextEmbedder(
api_key=Secret.from_env_var("OPENAI_API_KEY"),
model=embedding_model_name,
@@ -9,8 +9,6 @@ from haystack.utils.auth import Secret
from src.utils import load_env_vars
from ...trace import TraceGenerationInput, trace_generation
load_env_vars()
logging.getLogger("backoff").addHandler(logging.StreamHandler())
@@ -22,6 +20,7 @@ GENERATION_KWARGS = {
"temperature": 0,
"n": 1,
"max_tokens": MAX_TOKENS[MODEL_NAME] if MODEL_NAME in MAX_TOKENS else 4096,
"top_p": 1,
"response_format": {"type": "json_object"},
}
@@ -36,37 +35,10 @@ class CustomOpenAIGenerator(OpenAIGenerator):
)
@component
class TracedOpenAIGenerator(CustomOpenAIGenerator):
def _run(self, *args, **kwargs):
return super(TracedOpenAIGenerator, self).run(*args, **kwargs)
@component.output_types(replies=List[str], meta=List[Dict[str, Any]])
def run(
self,
trace_generation_input: TraceGenerationInput,
prompt: str,
generation_kwargs: Optional[Dict[str, Any]] = None,
):
return trace_generation(self._run)(
trace_generation_input=trace_generation_input,
prompt=prompt,
generation_kwargs=generation_kwargs,
)
def init_generator(
with_trace: bool = False,
model_name: str = MODEL_NAME,
generation_kwargs: Optional[Dict[str, Any]] = GENERATION_KWARGS,
):
if with_trace:
return TracedOpenAIGenerator(
api_key=Secret.from_env_var("OPENAI_API_KEY"),
model=model_name,
generation_kwargs=generation_kwargs,
)
return CustomOpenAIGenerator(
api_key=Secret.from_env_var("OPENAI_API_KEY"),
model=model_name,
@@ -21,17 +21,13 @@ class GenerationPostProcessor:
)
def run(self, replies: List[str]):
try:
cleaned_generation_result = json.loads(clean_generation_result(replies[0]))
cleaned_generation_result = json.loads(clean_generation_result(replies[0]))[
"results"
]
if isinstance(cleaned_generation_result, dict):
cleaned_generation_result = [cleaned_generation_result]
if cleaned_generation_result[0]["sql"] == "":
return {
"valid_generation_results": [],
"invalid_generation_results": [],
}
(
valid_generation_results,
invalid_generation_results,
@@ -1,41 +1,115 @@
from haystack.components.builders.prompt_builder import PromptBuilder
text_to_sql_user_prompt_template = """
You are a Trino SQL expert with exceptional logical thinking skills.
Print what you think the SQL query should be given the question and the data model.
This is vital to my career, I will become homeless if you make a mistake.
### INSTRUCTIONS ###
- If the question is complex enough, you can also answer complex SQL query that consists of a combination of JOINs, subqueries, and conditional filtering.
- Try not to use '*' to select all columns, please be specific what columns to choose from the table.
- If you can't construct the Trino SQL query, please answer with empty SQL string.
- If you can construct the Trino SQL query, please answer with the SQL query: ```sql ...```.
- Make sure the chosen "GROUP BY" conditions are correct given the selected columns.
- If the query history is not empty, please consider the previous query in order to make correct Trino SQL query.
### TASK ###
Given an input question, create a syntactically correct Trino SQL query to run and a short sentence within 10 words to summary the Trino SQL query
and return them as the answer to the input question.
Given a user query that is ambiguous in nature, your task is to interpret the query in various plausible ways and
generate three SQL statements that could potentially answer each interpreted version of the queries and within-10-words summary.
Provide three different interpretations and corresponding SQL queries that reflect these interpretations.
Ensure that your SQL queries are diverse, covering a range of possible meanings behind the ambiguous query.
### DATA MODELS ###
### EXAMPLES ###
Consider the structure of a generic database which includes common tables like users, orders, products, and transactions.
Here are the ambiguous user queries:
1. "Find the records of recent high-value transactions."
2. "Show me popular items that are not selling well."
3. "Retrieve user feedback on products from last month."
For each query, start by explaining the different ways the query can be interpreted. Then, provide SQL queries corresponding to each interpretation.
Your SQL statements should include SELECT statements, appropriate WHERE clauses to filter the results, and JOINs if necessary to combine information from different tables.
Remember to include ordering and limit clauses where relevant to address the 'recent', 'high-value', 'popular', and 'last month' aspects of the queries.
Example for the first query:
Interpretation 1: Recent high-value transactions are defined as transactions that occurred in the last 30 days with a value greater than $10,000.
SQL Query 1: SELECT * FROM transactions WHERE transaction_date >= NOW() - INTERVAL '30 days' AND value > 10000 ORDER BY transaction_date DESC;
SUMMARY 1: Recent high-value transactions.
Interpretation 2: High-value transactions are those in the top "10%" of all transactions in terms of value, and 'recent' is defined as the last 3 months.
SQL Query 2: WITH ranked_transactions AS (SELECT *, NTILE(10) OVER (ORDER BY value DESC) AS percentile_rank FROM transactions WHERE transaction_date >= NOW() - INTERVAL '3 months') SELECT * FROM ranked_transactions WHERE percentile_rank = 1 ORDER BY transaction_date DESC;
SUMMARY 2: Top 10% transactions last 3 months.
Interpretation 3: 'Recent' refers to the last week, and 'high-value' transactions are those above the average transaction value of the past week.
SQL Query 3: SELECT * FROM transactions WHERE transaction_date >= NOW() - INTERVAL '7 days' AND value > (SELECT AVG(value) FROM transactions WHERE transaction_date >= NOW() - INTERVAL '7 days') ORDER BY transaction_date DESC;
SUMMARY 3: Above-average transactions last week.
Proceed in a similar manner for the other queries.
### DATABASE SCHEMA ###
{% for document in documents %}
{{ document.content }}
{% endfor %}
### QUERY HISTORY ###
{{ history }}
### FINAL ANSWER FORMAT ###
The final answer must be the JSON format like following:
{
"results": [
{"sql": <SQL_QUERY_STRING_1>, "summary": <SUMMARY_STRING_1>},
{"sql": <SQL_QUERY_STRING2>, "summary": <SUMMARY_STRING_2>}
]
}
### NOTICE ###
- Only use the tables and columns mentioned in the database schema.
- If you think you can't generate a valid SQL query for a specific interpretation, you can skip that interpretation and provide the other ones.
- Make sure to map operators and operands correctly based on their data types.
### QUESTION ###
{{ query }}
"""
text_to_sql_with_followup_user_prompt_template = """
### TASK ###
Given the following user query and the history of the last query along with the generated SQL result,
generate appropriate SQL queries that match the user's current request.
Generate at most 3 SQL queries in order to interpret the user query in various plausible ways.
### EXAMPLES ###
Previous SQL Summary: "Users signed up this year."
Previous Generated SQL Query: "SELECT * FROM users WHERE sign_up_date >= '2023-01-01';"
Current User Query: "Who has made a purchase?"
Generated SQL Queries amd Summaries:
{
"results": [
{
"sql": "SELECT users.* FROM users JOIN purchases ON users.id = purchases.user_id WHERE users.sign_up_date >= '2023-01-01';",
"summary": "Users joined in 2023 with purchases."
},
{
"sql": "SELECT DISTINCT users.* FROM users INNER JOIN purchases ON users.id = purchases.user_id WHERE users.sign_up_date >= '2023-01-01';",
"summary": "Unique users with purchases since 2023."
}
]
}
### DATABASE SCHEMA ###
{% for document in documents %}
{{ document.content }}
{% endfor %}
### FINAL ANSWER FORMAT ###
The final answer must be the JSON format
The final answer must be the JSON format like following:
For a question that you can return the SQL query
{"sql": <SQL_QUERY_STRING>, "summary": <SUMMARY_STRING>}
{
"results": [
{"sql": <SQL_QUERY_STRING_1>, "summary": <SUMMARY_STRING_1>},
{"sql": <SQL_QUERY_STRING2>, "summary": <SUMMARY_STRING_2>}
]
}
For a question that you can't return the SQL query
{"sql": "", "summary": ""}
### NOTICE ###
- Only use the tables and columns mentioned in the database schema.
- If you think you can't generate a valid SQL query for a specific interpretation, you can skip that interpretation and provide the other ones.
- Make sure to map operators and operands correctly based on their data types.
### QUESTION ###
Previous SQL Summary: {{ history.summary }}
Previous Generated SQL Query: {{ history.sql }}
Current User Query: {{ query }}
Generated SQL Queries amd Summaries:
"""
sql_correction_user_prompt_template = """
@@ -43,20 +117,29 @@ You are a Trino SQL expert with exceptional logical thinking skills and debuggin
### TASK ###
Now you are given a list of syntactically incorrect Trino SQL queries and related error messages.
With given data models, please think step by step to correct these wrong Trino SQL quries.
With given database schema, please think step by step to correct these wrong Trino SQL quries.
### DATA MODELS ###
### DATABASE SCHEMA ###
{% for document in documents %}
{{ document.content }}
{% endfor %}
### QUESTION ###
{{ invalid_generation_results }}
### FINAL ANSWER FORMAT ###
The final answer must be a list of corrected SQL quries and its original corresponding summary in JSON format
{"sql": <CORRECTED_SQL_QUERY_STRING>, "summary": <ORIGINAL_SUMMARY_STRING>}
{
"results": [
{"sql": <CORRECTED_SQL_QUERY_STRING_1>, "summary": <ORIGINAL_SUMMARY_STRING_1>},
{"sql": <CORRECTED_SQL_QUERY_STRING_2>, "summary": <ORIGINAL_SUMMARY_STRING_2>}
]
}
### NOTICE ###
- Only use the tables and columns mentioned in the database schema.
- Make sure to map operators and operands correctly based on their data types.
### QUESTION ###
{{ invalid_generation_results }}
"""
@@ -64,5 +147,9 @@ def init_text_to_sql_prompt_builder():
return PromptBuilder(template=text_to_sql_user_prompt_template)
def init_text_to_sql_with_followup_prompt_builder():
return PromptBuilder(template=text_to_sql_with_followup_user_prompt_template)
def init_sql_correction_prompt_builder():
return PromptBuilder(template=sql_correction_user_prompt_template)
@@ -1,41 +1,7 @@
from typing import Any, Dict, List, Optional
from typing import Any
from haystack import Document, component
from haystack_integrations.components.retrievers.qdrant import QdrantEmbeddingRetriever
from ...trace import TraceSpanInput, trace_span
@component
class TracedQdrantEmbeddingRetriever(QdrantEmbeddingRetriever):
def _run(self, *args, **kwargs):
return super(TracedQdrantEmbeddingRetriever, self).run(*args, **kwargs)
@component.output_types(documents=List[Document])
def run(
self,
trace_span_input: TraceSpanInput,
query_embedding: List[float],
filters: Optional[Dict[str, Any]] = None,
top_k: Optional[int] = None,
scale_score: Optional[bool] = None,
return_embedding: Optional[bool] = None,
):
return trace_span(self._run)(
trace_span_input=trace_span_input,
query_embedding=query_embedding,
filters=filters,
top_k=top_k,
scale_score=scale_score,
return_embedding=return_embedding,
)
def init_retriever(document_store: Any, with_trace: bool = False, top_k: int = 3):
if with_trace:
return TracedQdrantEmbeddingRetriever(
document_store=document_store,
top_k=top_k,
)
def init_retriever(document_store: Any, top_k: int = 10):
return QdrantEmbeddingRetriever(document_store=document_store, top_k=top_k)
@@ -0,0 +1,66 @@
from typing import Any, List
from haystack import Document, Pipeline
from src.core.pipeline import BasicPipeline
from src.pipelines.ask.components.generator import (
init_generator,
)
from src.pipelines.ask.components.post_processors import init_generation_post_processor
from src.pipelines.ask.components.prompts import (
init_text_to_sql_with_followup_prompt_builder,
)
from src.utils import load_env_vars
from src.web.v1.services.ask import AskRequest
load_env_vars()
class FollowUpGeneration(BasicPipeline):
def __init__(
self,
generator: Any,
):
self._pipeline = Pipeline()
self._pipeline.add_component(
"text_to_sql_prompt_builder",
init_text_to_sql_with_followup_prompt_builder(),
)
self._pipeline.add_component("text_to_sql_generator", generator)
self._pipeline.add_component("post_processor", init_generation_post_processor())
self._pipeline.connect(
"text_to_sql_prompt_builder.prompt", "text_to_sql_generator.prompt"
)
self._pipeline.connect(
"text_to_sql_generator.replies", "post_processor.replies"
)
super().__init__(self._pipeline)
def run(
self,
query: str,
contexts: List[Document],
history: AskRequest.AskResponseDetails,
):
return self._pipeline.run(
{
"text_to_sql_prompt_builder": {
"query": query,
"documents": contexts,
"history": history,
},
}
)
if __name__ == "__main__":
followup_generation_pipeline = FollowUpGeneration(
generator=init_generator(),
)
print("generating followup_generation_pipeline.jpg to outputs/pipelines/ask...")
followup_generation_pipeline.draw(
"./outputs/pipelines/ask/followup_generation_pipeline.jpg"
)
@@ -1,40 +1,27 @@
import os
from typing import Any, List, Optional
from typing import Any, List
from haystack import Document, Pipeline
from src.core.pipeline import BasicPipeline
from src.pipelines.ask.components.generator import (
MODEL_NAME,
init_generator,
)
from src.pipelines.ask.components.generator import init_generator
from src.pipelines.ask.components.post_processors import init_generation_post_processor
from src.pipelines.ask.components.prompts import init_text_to_sql_prompt_builder
from src.utils import load_env_vars
from src.web.v1.services.ask import AskRequest
load_env_vars()
if with_trace := os.getenv("ENABLE_TRACE", default=False):
from src.pipelines.trace import (
TraceGenerationInput,
TraceInput,
langfuse,
)
class Generation(BasicPipeline):
def __init__(
self,
text_to_sql_generator: Any,
with_trace: bool = False,
generator: Any,
):
self._pipeline = Pipeline()
self._pipeline.add_component(
"text_to_sql_prompt_builder",
init_text_to_sql_prompt_builder(),
)
self._pipeline.add_component("text_to_sql_generator", text_to_sql_generator)
self._pipeline.add_component("text_to_sql_generator", generator)
self._pipeline.add_component("post_processor", init_generation_post_processor())
self._pipeline.connect(
@@ -43,70 +30,26 @@ class Generation(BasicPipeline):
self._pipeline.connect(
"text_to_sql_generator.replies", "post_processor.replies"
)
self.with_trace = with_trace
self.text_to_sql_prompt_builder = self._pipeline.get_component(
"text_to_sql_prompt_builder"
)
super().__init__(self._pipeline)
def run(
self,
query: str,
contexts: List[Document],
history: Optional[AskRequest.AskResponseDetails] = None,
user_id: Optional[str] = None,
):
if self.with_trace:
trace = langfuse.trace(
**TraceInput(
name="generation",
user_id=user_id,
).__dict__,
public=True,
)
result = self._pipeline.run(
{
"text_to_sql_prompt_builder": {
"query": query,
"documents": contexts,
"history": history,
},
"text_to_sql_generator": {
"trace_generation_input": TraceGenerationInput(
trace_id=trace.id,
name="generator",
input=self.text_to_sql_prompt_builder.run(
query=query,
documents=contexts,
history=history,
)["prompt"],
model=MODEL_NAME,
)
},
}
)
trace.update(input=query, output=result["text_to_sql_generator"])
else:
result = self._pipeline.run(
{
"text_to_sql_prompt_builder": {
"query": query,
"documents": contexts,
"history": history,
},
}
)
return result
return self._pipeline.run(
{
"text_to_sql_prompt_builder": {
"query": query,
"documents": contexts,
},
}
)
if __name__ == "__main__":
generation_pipeline = Generation(
text_to_sql_generator=init_generator(),
generator=init_generator(),
)
print("generating generation_pipeline.jpg to outputs/pipelines/ask...")
@@ -14,7 +14,7 @@ from src.pipelines.ask.components.embedder import (
EMBEDDING_MODEL_DIMENSION,
EMBEDDING_MODEL_NAME,
)
from src.utils import load_env_vars
from src.utils import generate_ddls_from_semantics, load_env_vars
load_env_vars()
@@ -88,11 +88,13 @@ class Indexing(BasicPipeline):
}
)
ddl_commands = generate_ddls_from_semantics(
semantics["models"],
semantics["relationships"],
)
embeddings = self._openai_client.embeddings.create(
input=[
json.dumps(data)
for data in semantics["models"] + semantics["relationships"]
],
input=ddl_commands,
model=self.embedding_model_name,
dimensions=self.embedding_model_dim,
)
@@ -101,12 +103,10 @@ class Indexing(BasicPipeline):
Document(
id=str(i),
meta={"id": str(i)},
content=json.dumps(data),
content=ddl_command,
embedding=embeddings.data[i].embedding,
)
for i, data in enumerate(
tqdm(semantics["models"] + semantics["relationships"])
)
for i, ddl_command in enumerate(tqdm(ddl_commands))
]
@@ -1,5 +1,4 @@
import os
from typing import Any, Optional
from typing import Any
from haystack import Pipeline
@@ -12,20 +11,12 @@ from src.utils import load_env_vars
load_env_vars()
if with_trace := os.getenv("ENABLE_TRACE", default=False):
from src.pipelines.trace import (
TraceInput,
TraceSpanInput,
langfuse,
)
class Retrieval(BasicPipeline):
def __init__(
self,
embedder: Any,
retriever: Any,
with_trace: bool = False,
):
self._pipeline = Pipeline()
self._pipeline.add_component("embedder", embedder)
@@ -35,51 +26,16 @@ class Retrieval(BasicPipeline):
self._pipeline.connect("embedder.embedding", "retriever.query_embedding")
self._pipeline.connect("retriever.documents", "post_processor.documents")
self.with_trace = with_trace
super().__init__(self._pipeline)
def run(self, query: str, user_id: Optional[str] = None):
if self.with_trace:
trace = langfuse.trace(
**TraceInput(
name="retrieval",
user_id=user_id,
).__dict__,
public=True,
)
result = self._pipeline.run(
{
"embedder": {
"trace_span_input": TraceSpanInput(
trace_id=trace.id,
name="text_embedder",
input=query,
),
"text": query,
},
"retriever": {
"trace_span_input": TraceSpanInput(
trace_id=trace.id,
name="retriever",
input="text_embedder.embedding",
),
},
}
)
trace.update(input=query, output=result["retriever"])
else:
result = self._pipeline.run(
{
"embedder": {
"text": query,
},
}
)
return result
def run(self, query: str):
return self._pipeline.run(
{
"embedder": {
"text": query,
},
}
)
if __name__ == "__main__":
@@ -15,16 +15,14 @@ from src.pipelines.ask.components.prompts import (
class SQLCorrection(BasicPipeline):
def __init__(
self,
sql_correction_generator: Any,
generator: Any,
):
self._pipeline = Pipeline()
self._pipeline.add_component(
"sql_correction_prompt_builder",
init_sql_correction_prompt_builder(),
)
self._pipeline.add_component(
"sql_correction_generator", sql_correction_generator
)
self._pipeline.add_component("sql_correction_generator", generator)
self._pipeline.add_component("post_processor", init_generation_post_processor())
self._pipeline.connect(
@@ -53,7 +51,7 @@ class SQLCorrection(BasicPipeline):
if __name__ == "__main__":
sql_correction_pipeline = SQLCorrection(
sql_correction_generator=init_generator(),
generator=init_generator(),
)
print("generating sql_correction_pipeline.jpg to outputs/pipelines/ask...")
@@ -8,8 +8,8 @@ from haystack.components.generators import OpenAIGenerator
from haystack.utils.auth import Secret
from src.utils import load_env_vars
from .prompts import init_sql_details_system_prompt_builder
from ...trace import TraceGenerationInput, trace_generation
load_env_vars()
@@ -37,40 +37,12 @@ class CustomOpenAIGenerator(OpenAIGenerator):
)
@component
class TracedOpenAIGenerator(CustomOpenAIGenerator):
def _run(self, *args, **kwargs):
return super(TracedOpenAIGenerator, self).run(*args, **kwargs)
@component.output_types(replies=List[str], meta=List[Dict[str, Any]])
def run(
self,
trace_generation_input: TraceGenerationInput,
prompt: str,
generation_kwargs: Optional[Dict[str, Any]] = None,
):
return trace_generation(self._run)(
trace_generation_input=trace_generation_input,
prompt=prompt,
generation_kwargs=generation_kwargs,
)
def init_generator(
with_trace: bool = False,
model_name: str = _MODEL_NAME,
generation_kwargs: Optional[Dict[str, Any]] = _GENERATION_KWARGS,
) -> Any:
system_prompt = init_sql_details_system_prompt_builder().run()["prompt"]
if with_trace:
return TracedOpenAIGenerator(
api_key=Secret.from_env_var("OPENAI_API_KEY"),
model=model_name,
generation_kwargs=generation_kwargs,
system_prompt=system_prompt,
)
return CustomOpenAIGenerator(
api_key=Secret.from_env_var("OPENAI_API_KEY"),
model=model_name,
@@ -16,7 +16,7 @@ load_env_vars()
@component
class GenerationPostProcessor:
@component.output_types(
post_processing_results=Optional[Dict[str, Any]],
results=Optional[Dict[str, Any]],
)
def run(self, replies: List[str], meta: List[Dict[str, Any]]) -> Dict[str, Any]:
generator = {
@@ -2,27 +2,74 @@ from haystack.components.builders.prompt_builder import PromptBuilder
sql_details_system_prompt_template = """
You are a Trino SQL expert with exceptional logical thinking skills.
Print what you think the SQL query means by giving 1 to 5 explainable steps to the user according to the complexity of SQL query.
If the SQL query is simple a select statement, you can just give one step to explain the SQL query; and vice versa.
This is vital to my career, I will become homeless if you make a mistake.
You are going to deconstruct a complex SQL query into one to five steps,
making it easier to understand. Each step has a SQL query part,
a summary explaining the purpose of that query,
and a CTE name to link the queries.
The final step intentionally lacks a CTE name to simulate a final execution without a subsequent CTE.
### TASK ###
Given an input SQL query, create two things:
1. a list of steps composed of syntactically and semantically correct Trino SQL query to run, a short sentence to summary the Trino SQL query and a cte_name to represent the Trino SQL query.
2. a short description describing the SQL query in a human-readable format.
3. there should be no CTEs in the SQL query in each step.
4. only the cte_name of the last step is empty.
### EXAMPLES ###
Example 1:
Original SQL Query:
SELECT product_id, SUM(sales) AS total_sales
FROM sales_data
GROUP BY product_id
HAVING SUM(sales) > 10000;
Results:
- Description: The breakdown simplifies the process of aggregating sales data by product and filtering for top-selling products.
- Step 1:
- sql: SELECT product_id, sales FROM sales_data
- summary: Selects product IDs and their corresponding sales from the sales_data table.
- cte_name: basic_sales_data
- Step 2:
- sql: SELECT product_id, SUM(sales) AS total_sales FROM basic_sales_data GROUP BY product_id
- summary: Aggregates sales by product, summing up sales for each product ID.
- cte_name: aggregated_sales
- Step 3:
- sql: SELECT product_id, total_sales FROM aggregated_sales WHERE total_sales > 10000
- summary: Filters the aggregated sales data to only include products whose total sales exceed 10,000.
- cte_name:
Example 2:
Original SQL Query:
SELECT product_id FROM sales_data
Results:
- Description: The breakdown simplifies the process of selecting product IDs from the sales_data table.
- Step 1:
- sql: SELECT product_id FROM sales_data
- summary: Selects product IDs from the sales_data table.
- cte_name:
### NOTICE ###
- Make sure to map operators and operands correctly based on their data types.
- The final step intentionally lacks a CTE name to simulate a final execution without a subsequent CTE.
- Only use the tables and columns mentioned in the original sql query.
### FINAL ANSWER FORMAT ###
The final answer must be a valid JSON format as follows:
The final answer must be a valid JSON format as following:
{
"description": <SHORT_SQL_QUERY_DESCRIPTION>,
"steps: [{
"sql": <SQL_QUERY_STRING>,
"summary": <SUMMARY_STRING>,
"cte_name": <CTE_NAME_STRING>
}] # list of steps
"steps: [
{
"sql": <SQL_QUERY_STRING_1>,
"summary": <SUMMARY_STRING_1>,
"cte_name": <CTE_NAME_STRING_1>
},
{
"sql": <SQL_QUERY_STRING_2>,
"summary": <SUMMARY_STRING_2>,
"cte_name": <CTE_NAME_STRING_2>
}
]
}
"""
@@ -1,10 +1,8 @@
import os
from typing import Any, Optional
from typing import Any
from haystack import Pipeline
from src.core.pipeline import BasicPipeline
from src.pipelines.ask.components.generator import MODEL_NAME
from src.pipelines.ask_details.components.generator import (
init_generator,
)
@@ -15,19 +13,11 @@ from src.utils import load_env_vars
load_env_vars()
if with_trace := os.getenv("ENABLE_TRACE", default=False):
from src.pipelines.trace import (
TraceGenerationInput,
TraceInput,
langfuse,
)
class Generation(BasicPipeline):
def __init__(
self,
generator: Any,
with_trace: bool = False,
):
self._pipeline = Pipeline()
self._pipeline.add_component("generator", generator)
@@ -35,43 +25,16 @@ class Generation(BasicPipeline):
self._pipeline.connect("generator.replies", "post_processor.replies")
self._pipeline.connect("generator.meta", "post_processor.meta")
self.with_trace = with_trace
super().__init__(self._pipeline)
def run(self, sql: str, user_id: Optional[str] = None):
if self.with_trace:
trace = langfuse.trace(
**TraceInput(
name="generation",
user_id=user_id,
).__dict__,
public=True,
)
result = self._pipeline.run(
{
"generator": {
"trace_generation_input": TraceGenerationInput(
trace_id=trace.id,
name="generator",
input=sql,
model=MODEL_NAME,
)
},
}
)
trace.update(input=sql, output=result["generator"])
return result
else:
return self._pipeline.run(
{
"generator": {
"prompt": sql,
},
}
)
def run(self, sql: str):
return self._pipeline.run(
{
"generator": {
"prompt": sql,
},
}
)
if __name__ == "__main__":
-96
View File
@@ -1,96 +0,0 @@
import os
from dataclasses import dataclass
from datetime import datetime
from typing import Any, Callable, Dict, List, Literal, Optional, Union
import pydantic
from langfuse import Langfuse
from langfuse.model import MapValue, ModelUsage, PromptClient
from src.utils import load_env_vars
load_env_vars()
if with_trace := os.getenv("ENABLE_TRACE", default=False):
langfuse = Langfuse(
public_key=os.getenv("LANGFUSE_PUBLIC_KEY"),
secret_key=os.getenv("LANGFUSE_SECRET_KEY"),
host="https://cloud.langfuse.com",
threads=os.cpu_count() // 2,
)
langfuse.auth_check()
@dataclass
class TraceInput:
id: Optional[str] = None
name: Optional[str] = None
user_id: Optional[str] = None
version: Optional[str] = None
input: Optional[Any] = None
output: Optional[Any] = None
metadata: Optional[Any] = None
tags: Optional[List[str]] = None
timestamp: Optional[datetime] = None
@dataclass
class TraceSpanInput:
id: Optional[str] = None
trace_id: Optional[str] = None
name: Optional[str] = None
start_time: Optional[datetime] = None
end_time: Optional[datetime] = None
metadata: Optional[Any] = None
input: Optional[Any] = None
output: Optional[Any] = None
level: Optional[Literal["DEBUG", "DEFAULT", "WARNING", "ERROR"]] = None
status_message: Optional[str] = None
parent_observation_id: Optional[str] = None
version: Optional[str] = None
@dataclass
class TraceGenerationInput:
id: Optional[str] = None
trace_id: Optional[str] = None
name: Optional[str] = None
start_time: Optional[datetime] = None
end_time: Optional[datetime] = None
metadata: Optional[Any] = None
level: Optional[Literal["DEBUG", "DEFAULT", "WARNING", "ERROR"]] = None
status_message: Optional[str] = None
parent_observation_id: Optional[str] = None
version: Optional[str] = None
completion_start_time: Optional[datetime] = None
completion_end_time: Optional[datetime] = None
model: Optional[str] = None
model_parameters: Optional[Dict[str, MapValue]] = None
input: Optional[Any] = None
output: Optional[Any] = None
usage: Optional[Union[pydantic.BaseModel, ModelUsage]] = None
prompt: Optional[PromptClient] = None
def trace_span(func: Callable):
def wrapper(*args, **kwargs):
span = langfuse.span(**kwargs["trace_span_input"].__dict__)
del kwargs["trace_span_input"]
results = func(*args, **kwargs)
span.end(output=results)
return results
return wrapper
def trace_generation(func: Callable):
def wrapper(*args, **kwargs):
generation = langfuse.generation(**kwargs["trace_generation_input"].__dict__)
del kwargs["trace_generation_input"]
results = func(*args, **kwargs)
generation.end(
output=results,
)
return results
return wrapper
+88 -8
View File
@@ -1,6 +1,7 @@
import json
import os
import re
from typing import Dict, List, Optional
from typing import Any, Dict, List, Optional
import requests
from dotenv import load_dotenv
@@ -18,10 +19,17 @@ def clean_generation_result(result: str) -> str:
.replace('"""', "")
.replace("'''", "")
.replace("```", "")
.replace(";", "")
)
def load_env_vars() -> str:
def _verify_env_vars() -> None:
"""
this is a temporary solution to verify that the required environment variables are set
"""
OpenAI().models.list()
load_dotenv(override=True)
if is_dev_env := os.getenv("ENV") and os.getenv("ENV").lower() == "dev":
@@ -33,11 +41,11 @@ def load_env_vars() -> str:
return "dev" if is_dev_env else "prod"
def _verify_env_vars() -> None:
"""
this is a temporary solution to verify that the required environment variables are set
"""
OpenAI().models.list()
def remove_limit_statement(sql: str) -> str:
pattern = r"\s*LIMIT\s+\d+(\s*;?\s*--.*|\s*;?\s*)$"
modified_sql = re.sub(pattern, "", sql, flags=re.IGNORECASE)
return modified_sql
def classify_invalid_generation_results(
@@ -51,7 +59,7 @@ def classify_invalid_generation_results(
response = requests.get(
f"{api_endpoint}/v1/mdl/preview",
json={
"sql": generation_result["sql"],
"sql": remove_limit_statement(generation_result["sql"]),
"limit": 1,
},
)
@@ -76,9 +84,81 @@ def check_if_sql_executable(
response = requests.get(
f"{api_endpoint}/v1/mdl/preview",
json={
"sql": sql,
"sql": remove_limit_statement(sql),
"limit": 1,
},
)
return True if response.status_code == 200 else False
def generate_ddls_from_semantics(
models: List[Dict[str, Any]],
relationships: List[Dict[str, Any]],
) -> List[str]:
ddl_commands = []
# A map to store model primary keys for foreign key relationships
primary_keys_map = {model["name"]: model["primaryKey"] for model in models}
for model in models:
table_name = model["name"]
columns_ddl = []
for column in model["columns"]:
if "relationship" not in column:
if column["properties"]:
comment = f"-- {json.dumps(column['properties'])}\n "
else:
comment = ""
column_name = column["name"]
column_type = column["type"]
column_ddl = f"{comment}{column_name} {column_type}"
# If column is a primary key
if column_name == model.get("primaryKey", ""):
column_ddl += " PRIMARY KEY"
columns_ddl.append(column_ddl)
# Add foreign key constraints based on relationships
for relationship in relationships:
if (
table_name == relationship["models"][0]
and relationship["joinType"].upper() == "MANY_TO_ONE"
):
related_table = relationship["models"][1]
fk_column = relationship["condition"].split(" = ")[0].split(".")[1]
fk_constraint = f"FOREIGN KEY ({fk_column}) REFERENCES {related_table}({primary_keys_map[related_table]})"
columns_ddl.append(fk_constraint)
elif (
table_name == relationship["models"][1]
and relationship["joinType"].upper() == "ONE_TO_MANY"
):
related_table = relationship["models"][0]
fk_column = relationship["condition"].split(" = ")[1].split(".")[1]
fk_constraint = f"FOREIGN KEY ({fk_column}) REFERENCES {related_table}({primary_keys_map[related_table]})"
columns_ddl.append(fk_constraint)
elif (
table_name in relationship["models"]
and relationship["joinType"].upper() == "ONE_TO_ONE"
):
index = relationship["models"].index(table_name)
related_table = [m for m in relationship["models"] if m != table_name][
0
]
fk_column = relationship["condition"].split(" = ")[index].split(".")[1]
fk_constraint = f"FOREIGN KEY ({fk_column}) REFERENCES {related_table}({primary_keys_map[related_table]})"
columns_ddl.append(fk_constraint)
if model["properties"]:
comment = f"\n/* {json.dumps(model['properties'])} */\n"
else:
comment = ""
create_table_ddl = (
f"{comment}CREATE TABLE {table_name} (\n "
+ ",\n ".join(columns_ddl)
+ "\n);"
)
ddl_commands.append(create_table_ddl)
return ddl_commands
+7 -7
View File
@@ -30,7 +30,7 @@ from src.web.v1.services.semantics import (
router = APIRouter()
@router.post("/semantics-descriptions/")
@router.post("/semantics-descriptions")
async def bulk_generate_description(
bulk_request: BulkGenerateDescriptionRequest,
) -> List[GenerateDescriptionResponse]:
@@ -40,7 +40,7 @@ async def bulk_generate_description(
]
@router.post("/semantics-preparations/")
@router.post("/semantics-preparations")
async def prepare_semantics(
prepare_semantics_request: SemanticsPreparationRequest,
background_tasks: BackgroundTasks,
@@ -52,7 +52,7 @@ async def prepare_semantics(
return SemanticsPreparationResponse(id=prepare_semantics_request.id)
@router.get("/semantics-preparations/{task_id}/status/")
@router.get("/semantics-preparations/{task_id}/status")
async def get_prepare_semantics_status(
task_id: str,
) -> SemanticsPreparationStatusResponse:
@@ -61,7 +61,7 @@ async def get_prepare_semantics_status(
)
@router.post("/asks/")
@router.post("/asks")
async def ask(
ask_request: AskRequest,
background_tasks: BackgroundTasks,
@@ -89,12 +89,12 @@ async def stop_ask(
return StopAskResponse(query_id=query_id)
@router.get("/asks/{query_id}/result/")
@router.get("/asks/{query_id}/result")
async def get_ask_result(query_id: str) -> AskResultResponse:
return container.ASK_SERVICE.get_ask_result(AskResultRequest(query_id=query_id))
@router.post("/ask-details/")
@router.post("/ask-details")
async def ask_details(
ask_details_request: AskDetailsRequest,
background_tasks: BackgroundTasks,
@@ -108,7 +108,7 @@ async def ask_details(
return AskDetailsResponse(query_id=query_id)
@router.get("/ask-details/{query_id}/result/")
@router.get("/ask-details/{query_id}/result")
async def get_ask_details_result(query_id: str) -> AskDetailsResultResponse:
return container.ASK_DETAILS_SERVICE.get_ask_details_result(
AskDetailsResultRequest(query_id=query_id)
+104 -65
View File
@@ -1,8 +1,11 @@
import logging
from typing import List, Literal, Optional
from haystack import Pipeline
from pydantic import BaseModel
logging.basicConfig(format="%(asctime)s %(message)s", level=logging.INFO)
# POST /v1/semantics-preparations
class SemanticsPreparationRequest(BaseModel):
@@ -21,6 +24,7 @@ class SemanticsPreparationStatusRequest(BaseModel):
class SemanticsPreparationStatusResponse(BaseModel):
status: Literal["indexing", "finished", "failed"]
error: Optional[str] = None
class SQLExplanation(BaseModel):
@@ -83,7 +87,9 @@ class AskResultResponse(BaseModel):
summary: str
class AskError(BaseModel):
code: Literal["MISLEADING_QUERY", "NO_RELEVANT_DATA", "NO_RELEVANT_SQL"]
code: Literal[
"MISLEADING_QUERY", "NO_RELEVANT_DATA", "NO_RELEVANT_SQL", "OTHERS"
]
message: str
status: Literal[
@@ -113,11 +119,12 @@ class AskService:
prepare_semantics_request.id
] = SemanticsPreparationStatusResponse(status="finished")
except Exception as e:
# TODO: log the error
print(f"Failed to prepare semantics: {e}")
self.prepare_semantics_statuses[
prepare_semantics_request.id
] = SemanticsPreparationStatusResponse(status="failed")
] = SemanticsPreparationStatusResponse(
status="failed",
error=f"Failed to prepare semantics: {e}",
)
def get_prepare_semantics_status(
self, prepare_semantics_status_request: SemanticsPreparationStatusRequest
@@ -134,76 +141,108 @@ class AskService:
self,
ask_request: AskRequest,
):
# ask status can be understanding, searching, generating, finished, failed, stopped
# we will need to handle business logic for each status
query_id = ask_request.query_id
try:
# ask status can be understanding, searching, generating, finished, failed, stopped
# we will need to handle business logic for each status
query_id = ask_request.query_id
if not self._is_stopped(query_id):
self.ask_results[query_id] = AskResultResponse(status="understanding")
if not self._is_stopped(query_id):
self.ask_results[query_id] = AskResultResponse(status="understanding")
if not self._is_stopped(query_id):
self.ask_results[query_id] = AskResultResponse(status="searching")
if not self._is_stopped(query_id):
self.ask_results[query_id] = AskResultResponse(status="searching")
retrieval_result = self._pipelines["retrieval"].run(
query=ask_request.query,
)
documents = retrieval_result["post_processor"]["documents"]
if not documents:
self.ask_results[query_id] = AskResultResponse(
status="failed",
error=AskResultResponse.AskError(
code="NO_RELEVANT_DATA",
message="No relevant data",
),
retrieval_result = self._pipelines["retrieval"].run(
query=ask_request.query,
)
return
documents = retrieval_result["post_processor"]["documents"]
if not self._is_stopped(query_id):
self.ask_results[query_id] = AskResultResponse(status="generating")
text_to_sql_generation_results = self._pipelines["generation"].run(
query=ask_request.query,
contexts=documents,
history=ask_request.history,
)
if not documents:
self.ask_results[query_id] = AskResultResponse(
status="failed",
error=AskResultResponse.AskError(
code="NO_RELEVANT_DATA",
message="No relevant data",
),
)
return
valid_generation_results = []
if text_to_sql_generation_results["post_processor"][
"valid_generation_results"
]:
valid_generation_results += text_to_sql_generation_results[
"post_processor"
]["valid_generation_results"]
if not self._is_stopped(query_id):
self.ask_results[query_id] = AskResultResponse(status="generating")
if ask_request.history:
text_to_sql_generation_results = self._pipelines[
"followup_generation"
].run(
query=ask_request.query,
contexts=documents,
history=ask_request.history,
)
else:
text_to_sql_generation_results = self._pipelines["generation"].run(
query=ask_request.query,
contexts=documents,
)
if text_to_sql_generation_results["post_processor"][
"invalid_generation_results"
]:
sql_correction_results = self._pipelines["sql_correction"].run(
contexts=documents,
invalid_generation_results=text_to_sql_generation_results[
"post_processor"
]["invalid_generation_results"],
)
valid_generation_results += sql_correction_results["post_processor"][
valid_generation_results = []
if text_to_sql_generation_results["post_processor"][
"valid_generation_results"
]
]:
valid_generation_results += text_to_sql_generation_results[
"post_processor"
]["valid_generation_results"]
if not valid_generation_results:
self.ask_results[query_id] = AskResultResponse(
status="failed",
error=AskResultResponse.AskError(
code="NO_RELEVANT_SQL",
message="No relevant SQL",
),
)
else:
self.ask_results[query_id] = AskResultResponse(
status="finished",
response=[
AskResultResponse.AskResult(**result)
for result in valid_generation_results
],
)
logging.debug("Documents:")
for document in documents:
logging.debug(f"score: {document.score}")
logging.debug(f"content: {document.content}")
logging.debug("Before sql correction:")
logging.debug(f"valid_generation_results: {valid_generation_results}")
if text_to_sql_generation_results["post_processor"][
"invalid_generation_results"
]:
sql_correction_results = self._pipelines["sql_correction"].run(
contexts=documents,
invalid_generation_results=text_to_sql_generation_results[
"post_processor"
]["invalid_generation_results"],
)
valid_generation_results += sql_correction_results[
"post_processor"
]["valid_generation_results"]
logging.debug(
f'sql_correction_results: {sql_correction_results["post_processor"]}'
)
logging.debug("After sql correction:")
logging.debug(f"valid_generation_results: {valid_generation_results}")
if not valid_generation_results:
self.ask_results[query_id] = AskResultResponse(
status="failed",
error=AskResultResponse.AskError(
code="NO_RELEVANT_SQL",
message="No relevant SQL",
),
)
else:
self.ask_results[query_id] = AskResultResponse(
status="finished",
response=[
AskResultResponse.AskResult(**result)
for result in valid_generation_results
],
)
except Exception as e:
self.ask_results[query_id] = AskResultResponse(
status="failed",
error=AskResultResponse.AskError(
code="OTHERS",
message=str(e),
),
)
def stop_ask(
self,
@@ -41,7 +41,7 @@ class AskDetailsResultResponse(BaseModel):
steps: List[SQLExplanation]
class AskDetailsError(BaseModel):
code: Literal["NO_RELEVANT_SQL"]
code: Literal["NO_RELEVANT_SQL", "OTHERS"]
message: str
status: Literal["understanding", "searching", "generating", "finished", "failed"]
@@ -58,40 +58,49 @@ class AskDetailsService:
self,
ask_details_request: AskDetailsRequest,
) -> AskDetailsResponse:
# ask details status can be understanding, searching, generating, finished, stopped
# we will need to handle business logic for each status
query_id = ask_details_request.query_id
try:
# ask details status can be understanding, searching, generating, finished, stopped
# we will need to handle business logic for each status
query_id = ask_details_request.query_id
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="understanding"
)
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="searching"
)
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="understanding"
)
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="searching"
)
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="generating"
)
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="generating"
)
generation_result = self._pipelines["generation"].run(
sql=ask_details_request.sql,
)
generation_result = self._pipelines["generation"].run(
sql=ask_details_request.sql,
)
ask_details_result = generation_result["post_processor"]["results"]
ask_details_result = generation_result["post_processor"]["results"]
if ask_details_result is None:
if ask_details_result is None:
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="failed",
error=AskDetailsResultResponse.AskDetailsError(
code="NO_RELEVANT_SQL",
message="No relevant SQL",
),
)
else:
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="finished",
response=AskDetailsResultResponse.AskDetailsResponseDetails(
**ask_details_result
),
)
except Exception as e:
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="failed",
error=AskDetailsResultResponse.AskDetailsError(
code="NO_RELEVANT_SQL",
message="No relevant SQL",
),
)
else:
self.ask_details_results[query_id] = AskDetailsResultResponse(
status="finished",
response=AskDetailsResultResponse.AskDetailsResponseDetails(
**ask_details_result
code="OTHERS",
message=str(e),
),
)
+31 -3
View File
@@ -7,11 +7,12 @@ from src.pipelines.ask.components.document_store import init_document_store
from src.pipelines.ask.components.embedder import init_embedder
from src.pipelines.ask.components.generator import init_generator
from src.pipelines.ask.components.retriever import init_retriever
from src.pipelines.ask.followup_generation_pipeline import FollowUpGeneration
from src.pipelines.ask.generation_pipeline import Generation
from src.pipelines.ask.indexing_pipeline import Indexing
from src.pipelines.ask.retrieval_pipeline import Retrieval
from src.pipelines.ask.sql_correction_pipeline import SQLCorrection
from src.web.v1.services.ask import AskResultResponse
from src.web.v1.services.ask import AskRequest, AskResultResponse, SQLExplanation
GLOBAL_DATA = {
"contexts": None,
@@ -57,7 +58,7 @@ def test_retrieval_pipeline(document_store: Any):
def test_generation_pipeline():
generation_pipeline = Generation(
text_to_sql_generator=init_generator(),
generator=init_generator(),
)
generation_result = generation_pipeline.run(
"How many authors are there?",
@@ -69,9 +70,36 @@ def test_generation_pipeline():
)
def test_followup_generation_pipeline():
generation_pipeline = FollowUpGeneration(
generator=init_generator(),
)
generation_result = generation_pipeline.run(
"What are names of the books?",
contexts=GLOBAL_DATA["contexts"],
history=AskRequest.AskResponseDetails(
sql="SELECT COUNT(*) FROM book",
summary="Retrieve the number of books",
steps=[
SQLExplanation(
sql="SELECT COUNT(*) FROM book",
summary="Retrieve the number of books",
cte_name="",
)
],
),
)
print(generation_result)
assert AskResultResponse.AskResult(
**generation_result["post_processor"]["valid_generation_results"][0]
)
def test_sql_correction_pipeline():
sql_correction_pipeline = SQLCorrection(
sql_correction_generator=init_generator(),
generator=init_generator(),
)
sql_correction_result = sql_correction_pipeline.run(
+2 -2
View File
@@ -39,10 +39,10 @@ def ask_service():
retriever=retriever,
),
"generation": generation_pipeline.Generation(
text_to_sql_generator=text_to_sql_generator,
generator=text_to_sql_generator,
),
"sql_correction": sql_correction_pipeline.SQLCorrection(
sql_correction_generator=sql_correction_generator,
generator=sql_correction_generator,
),
}
)