mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
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:
co-authored by
qa
Aster Sun
imAsterSun
Pao Sheng
parent
b677dd1f77
commit
9507838991
+1
-1
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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="🚨",
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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__":
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user