feat(wren-ai-service): support multi compnents(llm, embedder, engine, document store) configuration (#682)

* chore: add pyyaml and update dependencies

* chore: add pydantic-settings

* feat: using pipeline component to impl specified component for each pipes

* chore: move pipeline component class to core pipeline module

* feat: service metadata includes the model metadata for each pipe

* feat: generate component for each pipe from config

* feat: initialize providers and dispatch them to specified pipe component

* feat: process diff kind of entry from yaml file

* feat: initialize provider by vender name

* feat: ignoring unused parameter in provider initializer

* chore: rename the property vender to provider to unify the name in code and config

* feat: unify the name in config and code

* feat: default spec of config example file

* chore: remove force deploy in service

* fix: modify some behavior to avoid service starting failed

* feat: decrease the dependency with haystack Secret

* chore: move init provider function to provider module

* feat: compatibility if config file not exist

* chore: modify the pattern of ignoring config yaml

* chore: update poetry.lock after merging the new commits

* fix: breaking changes of langfuse metadata tracing

* chore: rename initializer parameter to match the properties in config file converter

* chore: add the default kwargs for llm model

* chore: correct the wrong embedder model in example

* update write engine env vars

* add config.yaml in reload_includes

* remove unused code

* fix engine support for multi llm in demo

---------

Co-authored-by: ChihYu Yeh <chihyu.jimmy.yeh@gmail.com>
This commit is contained in:
Pao Sheng
2024-10-01 14:08:24 +08:00
committed by GitHub
co-authored by ChihYu Yeh
parent 62f5a2c757
commit 7f4a0bebc7
38 changed files with 612 additions and 343 deletions
+1
View File
@@ -1,5 +1,6 @@
# wren-ai-service
wren-ai-service/.env.*
wren-ai-service/config.yaml*
!wren-ai-service/.env.*.example
!wren-ai-service/src/eval/wren-engine/.env
wren-ai-service/src/eval/wren-engine/**/config.properties
+105
View File
@@ -0,0 +1,105 @@
type: llm
provider: openai_llm
models:
- model: gpt-4o-mini
kwargs:
{
"temperature": 0,
"n": 1,
"max_tokens": 4096,
"response_format": { "type": "json_object" },
}
api_key: sk-xxx
api_base: https://api.openai.com/v1
---
type: llm
provider: azure_openai_llm
models:
- model: gpt-4o
kwargs:
{
"temperature": 0,
"n": 1,
"max_tokens": 4096,
"response_format": { "type": "json_object" },
}
api_key: sk-xxx
api_base: https://api.openai.com/v1
api_version: "2024-05-13"
---
type: embedder
provider: openai_embedder
models:
- model: text-embedding-3-large
dimension: 3072
api_key: sk-xxx
api_base: https://api.openai.com/v1
timeout: 120
---
type: engine
provider: wren_ui
endpoint: http://localhost:3000
---
type: engine
provider: wren_ibis
endpoint: http://localhost:8000
source: bigquery
manifest: "" # base64 encoded string of the MDL
connection_info: "" # base64 encoded string of the connection info
---
type: engine
provider: wren_engine
endpoint: http://localhost:8080
manifest: ""
---
type: document_store
provider: qdrant
location: http://localhost:6333
api_key: ""
embedding_model_dim: 3072
timeout: 120
---
type: pipeline
pipes:
- name: indexing
embedder: openai_embedder.text-embedding-3-large
document_store: qdrant
- name: retrieval
llm: openai_llm.gpt-4o-mini
embedder: openai_embedder.text-embedding-3-large
document_store: qdrant
- name: historical_question
embedder: openai_embedder.text-embedding-3-large
document_store: qdrant
- name: sql_generation
llm: openai_llm.gpt-4o-mini
engine: wren_ui
- name: sql_correction
llm: openai_llm.gpt-4o-mini
engine: wren_ui
- name: followup_sql_generation
llm: openai_llm.gpt-4o-mini
engine: wren_ui
- name: sql_summary
llm: openai_llm.gpt-4o-mini
- name: sql_answer
llm: openai_llm.gpt-4o-mini
engine: wren_ui
- name: sql_breakdown
llm: openai_llm.gpt-4o-mini
engine: wren_ui
- name: sql_expansion
llm: openai_llm.gpt-4o-mini
engine: wren_ui
- name: sql_explanation
llm: openai_llm.gpt-4o-mini
- name: sql_regeneration
llm: openai_llm.gpt-4o-mini
engine: wren_ui
-26
View File
@@ -5,26 +5,19 @@ import orjson
import streamlit as st
from utils import (
DATA_SOURCES,
LLM_MODELS,
ask,
ask_details,
get_default_llm_model,
get_mdl_json,
prepare_semantics,
rerun_wren_engine,
save_mdl_json_file,
show_asks_details_results,
show_asks_results,
update_llm,
)
st.set_page_config(layout="wide")
st.title("Wren AI LLM Service Demo")
llm_model = get_default_llm_model(LLM_MODELS)
if "chosen_llm_model" not in st.session_state:
st.session_state["chosen_llm_model"] = llm_model
if "deployment_id" not in st.session_state:
st.session_state["deployment_id"] = str(uuid.uuid4())
if "chosen_dataset" not in st.session_state:
@@ -69,25 +62,6 @@ def onchange_demo_dataset():
st.session_state["chosen_dataset"] = st.session_state["choose_demo_dataset"]
def onchange_llm_model():
if (
st.session_state["llm_model_selectbox"]
and st.session_state["chosen_llm_model"]
!= st.session_state["llm_model_selectbox"]
):
st.session_state["chosen_llm_model"] = st.session_state["llm_model_selectbox"]
update_llm(st.session_state["chosen_llm_model"], st.session_state["mdl_json"])
st.selectbox(
"Select an OpenAI LLM model",
LLM_MODELS,
index=LLM_MODELS.index(llm_model),
key="llm_model_selectbox",
on_change=onchange_llm_model,
)
with st.sidebar:
st.markdown("## Deploy MDL Model")
uploaded_file = st.file_uploader(
+59 -61
View File
@@ -12,6 +12,7 @@ import requests
import sqlglot
import sqlparse
import streamlit as st
import yaml
from dotenv import load_dotenv
WREN_AI_SERVICE_BASE_URL = "http://localhost:5556"
@@ -19,7 +20,6 @@ WREN_ENGINE_API_URL = "http://localhost:8080"
WREN_IBIS_API_URL = "http://localhost:8000"
POLLING_INTERVAL = 0.5
DATA_SOURCES = ["duckdb", "bigquery", "postgres"]
LLM_MODELS = ["gpt-4o-mini", "gpt-4o", "gpt-4-turbo", "gpt-3.5-turbo"]
load_dotenv()
@@ -74,41 +74,22 @@ def rerun_wren_engine(mdl_json: Dict, dataset_type: str, dataset: str):
)
_prepare_duckdb(dataset)
# replace the values of WREN_ENGINE_xxx to ../.env.dev
with open(".env.dev", "r") as f:
lines = f.readlines()
for i, line in enumerate(lines):
if line.startswith("ENGINE"):
lines[i] = "ENGINE=wren_engine\n"
elif line.startswith("WREN_ENGINE_MANIFEST"):
lines[i] = f"WREN_ENGINE_MANIFEST={MANIFEST}\n"
with open(".env.dev", "w") as f:
f.writelines(lines)
_replace_wren_engine_env_variables("wren_engine", {"manifest": MANIFEST})
else:
WREN_IBIS_CONNECTION_INFO = base64.b64encode(
orjson.dumps(get_connection_info(dataset_type))
).decode()
# replace the values of WREN_IBIS_xxx to ../.env.dev
with open(".env.dev", "r") as f:
lines = f.readlines()
for i, line in enumerate(lines):
if line.startswith("ENGINE"):
lines[i] = "ENGINE=wren_ibis\n"
elif line.startswith("WREN_IBIS_SOURCE"):
lines[i] = f"WREN_IBIS_SOURCE={SOURCE}\n"
elif line.startswith("WREN_IBIS_MANIFEST"):
lines[i] = f"WREN_IBIS_MANIFEST={MANIFEST}\n"
elif (
line.startswith("WREN_IBIS_CONNECTION_INFO")
and dataset_type != "duckdb"
):
lines[
i
] = f"WREN_IBIS_CONNECTION_INFO={WREN_IBIS_CONNECTION_INFO}\n"
with open(".env.dev", "w") as f:
f.writelines(lines)
_replace_wren_engine_env_variables(
"wren_ibis",
{
"manifest": MANIFEST,
"source": SOURCE,
"connection_info": WREN_IBIS_CONNECTION_INFO
if dataset_type != "duckdb"
else "",
},
)
# wait for wren-ai-service to restart
time.sleep(5)
@@ -587,6 +568,53 @@ CREATE TABLE reviews AS FROM read_csv('https://wrenai-public.s3.amazonaws.com/de
assert response.status_code == 200, response.text
def _replace_wren_engine_env_variables(engine_type: str, data: dict):
assert engine_type in ("wren_engine", "wren_ibis")
if not Path("config.yaml").exists():
if engine_type == "wren_engine":
with open(".env.dev", "r") as f:
lines = f.readlines()
for i, line in enumerate(lines):
if line.startswith("ENGINE"):
lines[i] = "ENGINE=wren_engine\n"
elif line.startswith("WREN_ENGINE_MANIFEST"):
lines[i] = f"WREN_ENGINE_MANIFEST={data['manifest']}\n"
else:
with open(".env.dev", "r") as f:
lines = f.readlines()
for i, line in enumerate(lines):
if line.startswith("ENGINE"):
lines[i] = "ENGINE=wren_ibis\n"
elif line.startswith("WREN_IBIS_SOURCE"):
lines[i] = f"WREN_IBIS_SOURCE={data['source']}\n"
elif line.startswith("WREN_IBIS_MANIFEST"):
lines[i] = f"WREN_IBIS_MANIFEST={data['manifest']}\n"
elif line.startswith("WREN_IBIS_CONNECTION_INFO"):
lines[
i
] = f"WREN_IBIS_CONNECTION_INFO={data['connection_info']}\n"
with open(".env.dev", "w") as f:
f.writelines(lines)
else:
with open("config.yaml", "r+") as f:
configs = list(yaml.safe_load_all(f))
for config in configs:
if config["type"] == "engine" and config["provider"] == engine_type:
for key, value in data.items():
config[key] = value
if "pipes" in config:
for i, pipe in enumerate(config["pipes"]):
if "engine" in pipe:
config["pipes"][i]["engine"] = engine_type
f.seek(0)
yaml.safe_dump_all(configs, f, default_flow_style=False)
def prepare_semantics(mdl_json: dict):
semantics_preparation_response = requests.post(
f"{WREN_AI_SERVICE_BASE_URL}/v1/semantics-preparations",
@@ -885,33 +913,3 @@ def show_sql_regeneration_results_dialog(
language="sql",
)
sqls_with_cte.append(f"{step['cte_name']} AS ( {step['sql']} )")
@st.cache_data
def update_llm(chosen_llm_model: str, mdl_json: dict):
with open(".env.dev", "r") as f:
lines = f.readlines()
for i, line in enumerate(lines):
if line.startswith("GENERATION_MODEL"):
lines[i] = f"GENERATION_MODEL={chosen_llm_model}\n"
break
with open(".env.dev", "w") as f:
f.writelines(lines)
# wait for wren-ai-service to restart
time.sleep(5)
prepare_semantics(mdl_json)
def get_default_llm_model(llm_models: list[str]):
with open(".env.dev", "r") as f:
lines = f.readlines()
for line in lines:
if line.startswith("GENERATION_MODEL"):
llm_model = line.split("=")[1].strip()
break
assert llm_model in llm_models
return llm_model
+2 -1
View File
@@ -15,6 +15,7 @@ from tomlkit import document, dumps
sys.path.append(f"{Path().parent.resolve()}")
import eval.pipelines as pipelines
import src.providers as provider
import src.utils as utils
from eval.utils import parse_toml
from src.core.engine import EngineConfig
@@ -105,7 +106,7 @@ def init_providers(mdl: dict) -> dict:
},
)
providers = utils.init_providers(engine_config=engine_config)
providers = provider.init_providers(engine_config=engine_config)
return {
"llm_provider": providers[0],
"embedder_provider": providers[1],
+20 -62
View File
@@ -1603,78 +1603,36 @@ optional = false
python-versions = ">=3.7"
files = [
{file = "greenlet-3.1.1-cp310-cp310-macosx_11_0_universal2.whl", hash = "sha256:0bbae94a29c9e5c7e4a2b7f0aae5c17e8e90acbfd3bf6270eeba60c39fce3563"},
{file = "greenlet-3.1.1-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:0fde093fb93f35ca72a556cf72c92ea3ebfda3d79fc35bb19fbe685853869a83"},
{file = "greenlet-3.1.1-cp310-cp310-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:36b89d13c49216cadb828db8dfa6ce86bbbc476a82d3a6c397f0efae0525bdd0"},
{file = "greenlet-3.1.1-cp310-cp310-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:94b6150a85e1b33b40b1464a3f9988dcc5251d6ed06842abff82e42632fac120"},
{file = "greenlet-3.1.1-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:93147c513fac16385d1036b7e5b102c7fbbdb163d556b791f0f11eada7ba65dc"},
{file = "greenlet-3.1.1-cp310-cp310-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:da7a9bff22ce038e19bf62c4dd1ec8391062878710ded0a845bcf47cc0200617"},
{file = "greenlet-3.1.1-cp310-cp310-musllinux_1_1_aarch64.whl", hash = "sha256:b2795058c23988728eec1f36a4e5e4ebad22f8320c85f3587b539b9ac84128d7"},
{file = "greenlet-3.1.1-cp310-cp310-musllinux_1_1_x86_64.whl", hash = "sha256:ed10eac5830befbdd0c32f83e8aa6288361597550ba669b04c48f0f9a2c843c6"},
{file = "greenlet-3.1.1-cp310-cp310-win_amd64.whl", hash = "sha256:77c386de38a60d1dfb8e55b8c1101d68c79dfdd25c7095d51fec2dd800892b80"},
{file = "greenlet-3.1.1-cp311-cp311-macosx_11_0_universal2.whl", hash = "sha256:e4d333e558953648ca09d64f13e6d8f0523fa705f51cae3f03b5983489958c70"},
{file = "greenlet-3.1.1-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:09fc016b73c94e98e29af67ab7b9a879c307c6731a2c9da0db5a7d9b7edd1159"},
{file = "greenlet-3.1.1-cp311-cp311-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:d5e975ca70269d66d17dd995dafc06f1b06e8cb1ec1e9ed54c1d1e4a7c4cf26e"},
{file = "greenlet-3.1.1-cp311-cp311-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:3b2813dc3de8c1ee3f924e4d4227999285fd335d1bcc0d2be6dc3f1f6a318ec1"},
{file = "greenlet-3.1.1-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:e347b3bfcf985a05e8c0b7d462ba6f15b1ee1c909e2dcad795e49e91b152c383"},
{file = "greenlet-3.1.1-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9e8f8c9cb53cdac7ba9793c276acd90168f416b9ce36799b9b885790f8ad6c0a"},
{file = "greenlet-3.1.1-cp311-cp311-musllinux_1_1_aarch64.whl", hash = "sha256:62ee94988d6b4722ce0028644418d93a52429e977d742ca2ccbe1c4f4a792511"},
{file = "greenlet-3.1.1-cp311-cp311-musllinux_1_1_x86_64.whl", hash = "sha256:1776fd7f989fc6b8d8c8cb8da1f6b82c5814957264d1f6cf818d475ec2bf6395"},
{file = "greenlet-3.1.1-cp311-cp311-win_amd64.whl", hash = "sha256:48ca08c771c268a768087b408658e216133aecd835c0ded47ce955381105ba39"},
{file = "greenlet-3.1.1-cp312-cp312-macosx_11_0_universal2.whl", hash = "sha256:4afe7ea89de619adc868e087b4d2359282058479d7cfb94970adf4b55284574d"},
{file = "greenlet-3.1.1-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:f406b22b7c9a9b4f8aa9d2ab13d6ae0ac3e85c9a809bd590ad53fed2bf70dc79"},
{file = "greenlet-3.1.1-cp312-cp312-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:c3a701fe5a9695b238503ce5bbe8218e03c3bcccf7e204e455e7462d770268aa"},
{file = "greenlet-3.1.1-cp312-cp312-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:2846930c65b47d70b9d178e89c7e1a69c95c1f68ea5aa0a58646b7a96df12441"},
{file = "greenlet-3.1.1-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:99cfaa2110534e2cf3ba31a7abcac9d328d1d9f1b95beede58294a60348fba36"},
{file = "greenlet-3.1.1-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1443279c19fca463fc33e65ef2a935a5b09bb90f978beab37729e1c3c6c25fe9"},
{file = "greenlet-3.1.1-cp312-cp312-musllinux_1_1_aarch64.whl", hash = "sha256:b7cede291382a78f7bb5f04a529cb18e068dd29e0fb27376074b6d0317bf4dd0"},
{file = "greenlet-3.1.1-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:23f20bb60ae298d7d8656c6ec6db134bca379ecefadb0b19ce6f19d1f232a942"},
{file = "greenlet-3.1.1-cp312-cp312-win_amd64.whl", hash = "sha256:7124e16b4c55d417577c2077be379514321916d5790fa287c9ed6f23bd2ffd01"},
{file = "greenlet-3.1.1-cp313-cp313-macosx_11_0_universal2.whl", hash = "sha256:05175c27cb459dcfc05d026c4232f9de8913ed006d42713cb8a5137bd49375f1"},
{file = "greenlet-3.1.1-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:935e943ec47c4afab8965954bf49bfa639c05d4ccf9ef6e924188f762145c0ff"},
{file = "greenlet-3.1.1-cp313-cp313-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:667a9706c970cb552ede35aee17339a18e8f2a87a51fba2ed39ceeeb1004798a"},
{file = "greenlet-3.1.1-cp313-cp313-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:b8a678974d1f3aa55f6cc34dc480169d58f2e6d8958895d68845fa4ab566509e"},
{file = "greenlet-3.1.1-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:efc0f674aa41b92da8c49e0346318c6075d734994c3c4e4430b1c3f853e498e4"},
{file = "greenlet-3.1.1-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0153404a4bb921f0ff1abeb5ce8a5131da56b953eda6e14b88dc6bbc04d2049e"},
{file = "greenlet-3.1.1-cp313-cp313-musllinux_1_1_aarch64.whl", hash = "sha256:275f72decf9932639c1c6dd1013a1bc266438eb32710016a1c742df5da6e60a1"},
{file = "greenlet-3.1.1-cp313-cp313-musllinux_1_1_x86_64.whl", hash = "sha256:c4aab7f6381f38a4b42f269057aee279ab0fc7bf2e929e3d4abfae97b682a12c"},
{file = "greenlet-3.1.1-cp313-cp313-win_amd64.whl", hash = "sha256:b42703b1cf69f2aa1df7d1030b9d77d3e584a70755674d60e710f0af570f3761"},
{file = "greenlet-3.1.1-cp313-cp313t-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:f1695e76146579f8c06c1509c7ce4dfe0706f49c6831a817ac04eebb2fd02011"},
{file = "greenlet-3.1.1-cp313-cp313t-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:7876452af029456b3f3549b696bb36a06db7c90747740c5302f74a9e9fa14b13"},
{file = "greenlet-3.1.1-cp313-cp313t-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:4ead44c85f8ab905852d3de8d86f6f8baf77109f9da589cb4fa142bd3b57b475"},
{file = "greenlet-3.1.1-cp313-cp313t-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:8320f64b777d00dd7ccdade271eaf0cad6636343293a25074cc5566160e4de7b"},
{file = "greenlet-3.1.1-cp313-cp313t-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:6510bf84a6b643dabba74d3049ead221257603a253d0a9873f55f6a59a65f822"},
{file = "greenlet-3.1.1-cp313-cp313t-musllinux_1_1_aarch64.whl", hash = "sha256:04b013dc07c96f83134b1e99888e7a79979f1a247e2a9f59697fa14b5862ed01"},
{file = "greenlet-3.1.1-cp313-cp313t-musllinux_1_1_x86_64.whl", hash = "sha256:411f015496fec93c1c8cd4e5238da364e1da7a124bcb293f085bf2860c32c6f6"},
{file = "greenlet-3.1.1-cp37-cp37m-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:47da355d8687fd65240c364c90a31569a133b7b60de111c255ef5b606f2ae291"},
{file = "greenlet-3.1.1-cp37-cp37m-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:98884ecf2ffb7d7fe6bd517e8eb99d31ff7855a840fa6d0d63cd07c037f6a981"},
{file = "greenlet-3.1.1-cp37-cp37m-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:f1d4aeb8891338e60d1ab6127af1fe45def5259def8094b9c7e34690c8858803"},
{file = "greenlet-3.1.1-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:db32b5348615a04b82240cc67983cb315309e88d444a288934ee6ceaebcad6cc"},
{file = "greenlet-3.1.1-cp37-cp37m-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:dcc62f31eae24de7f8dce72134c8651c58000d3b1868e01392baea7c32c247de"},
{file = "greenlet-3.1.1-cp37-cp37m-musllinux_1_1_aarch64.whl", hash = "sha256:1d3755bcb2e02de341c55b4fca7a745a24a9e7212ac953f6b3a48d117d7257aa"},
{file = "greenlet-3.1.1-cp37-cp37m-musllinux_1_1_x86_64.whl", hash = "sha256:b8da394b34370874b4572676f36acabac172602abf054cbc4ac910219f3340af"},
{file = "greenlet-3.1.1-cp37-cp37m-win32.whl", hash = "sha256:a0dfc6c143b519113354e780a50381508139b07d2177cb6ad6a08278ec655798"},
{file = "greenlet-3.1.1-cp37-cp37m-win_amd64.whl", hash = "sha256:54558ea205654b50c438029505def3834e80f0869a70fb15b871c29b4575ddef"},
{file = "greenlet-3.1.1-cp38-cp38-macosx_11_0_universal2.whl", hash = "sha256:346bed03fe47414091be4ad44786d1bd8bef0c3fcad6ed3dee074a032ab408a9"},
{file = "greenlet-3.1.1-cp38-cp38-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:dfc59d69fc48664bc693842bd57acfdd490acafda1ab52c7836e3fc75c90a111"},
{file = "greenlet-3.1.1-cp38-cp38-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:d21e10da6ec19b457b82636209cbe2331ff4306b54d06fa04b7c138ba18c8a81"},
{file = "greenlet-3.1.1-cp38-cp38-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:37b9de5a96111fc15418819ab4c4432e4f3c2ede61e660b1e33971eba26ef9ba"},
{file = "greenlet-3.1.1-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:6ef9ea3f137e5711f0dbe5f9263e8c009b7069d8a1acea822bd5e9dae0ae49c8"},
{file = "greenlet-3.1.1-cp38-cp38-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:85f3ff71e2e60bd4b4932a043fbbe0f499e263c628390b285cb599154a3b03b1"},
{file = "greenlet-3.1.1-cp38-cp38-musllinux_1_1_aarch64.whl", hash = "sha256:95ffcf719966dd7c453f908e208e14cde192e09fde6c7186c8f1896ef778d8cd"},
{file = "greenlet-3.1.1-cp38-cp38-musllinux_1_1_x86_64.whl", hash = "sha256:03a088b9de532cbfe2ba2034b2b85e82df37874681e8c470d6fb2f8c04d7e4b7"},
{file = "greenlet-3.1.1-cp38-cp38-win32.whl", hash = "sha256:8b8b36671f10ba80e159378df9c4f15c14098c4fd73a36b9ad715f057272fbef"},
{file = "greenlet-3.1.1-cp38-cp38-win_amd64.whl", hash = "sha256:7017b2be767b9d43cc31416aba48aab0d2309ee31b4dbf10a1d38fb7972bdf9d"},
{file = "greenlet-3.1.1-cp39-cp39-macosx_11_0_universal2.whl", hash = "sha256:396979749bd95f018296af156201d6211240e7a23090f50a8d5d18c370084dc3"},
{file = "greenlet-3.1.1-cp39-cp39-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ca9d0ff5ad43e785350894d97e13633a66e2b50000e8a183a50a88d834752d42"},
{file = "greenlet-3.1.1-cp39-cp39-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:f6ff3b14f2df4c41660a7dec01045a045653998784bf8cfcb5a525bdffffbc8f"},
{file = "greenlet-3.1.1-cp39-cp39-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:94ebba31df2aa506d7b14866fed00ac141a867e63143fe5bca82a8e503b36437"},
{file = "greenlet-3.1.1-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:73aaad12ac0ff500f62cebed98d8789198ea0e6f233421059fa68a5aa7220145"},
{file = "greenlet-3.1.1-cp39-cp39-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:63e4844797b975b9af3a3fb8f7866ff08775f5426925e1e0bbcfe7932059a12c"},
{file = "greenlet-3.1.1-cp39-cp39-musllinux_1_1_aarch64.whl", hash = "sha256:7939aa3ca7d2a1593596e7ac6d59391ff30281ef280d8632fa03d81f7c5f955e"},
{file = "greenlet-3.1.1-cp39-cp39-musllinux_1_1_x86_64.whl", hash = "sha256:d0028e725ee18175c6e422797c407874da24381ce0690d6b9396c204c7f7276e"},
{file = "greenlet-3.1.1-cp39-cp39-win32.whl", hash = "sha256:5e06afd14cbaf9e00899fae69b24a32f2196c19de08fcb9f4779dd4f004e5e7c"},
{file = "greenlet-3.1.1-cp39-cp39-win_amd64.whl", hash = "sha256:3319aa75e0e0639bc15ff54ca327e8dc7a6fe404003496e3c6925cd3142e0e22"},
{file = "greenlet-3.1.1.tar.gz", hash = "sha256:4ce3ac6cdb6adf7946475d7ef31777c26d94bccc377e070a7986bd2d5c515467"},
]
[package.extras]
@@ -1857,13 +1815,13 @@ typing-extensions = ">=4.7"
[[package]]
name = "haystack-experimental"
version = "0.1.1"
version = "0.2.0"
description = "Experimental components and features for the Haystack LLM framework."
optional = false
python-versions = ">=3.8"
files = [
{file = "haystack_experimental-0.1.1-py3-none-any.whl", hash = "sha256:c11fa89c0754dafd3ec25220c02e459d784acda7a9c047e3297e1b11f35ebb06"},
{file = "haystack_experimental-0.1.1.tar.gz", hash = "sha256:64cc58d54b9cf9a2957616e4cb44a10c12e36ab640ebe0c38e7cc70300cb6726"},
{file = "haystack_experimental-0.2.0-py3-none-any.whl", hash = "sha256:a4601d12d824a6fc2f4f6fb3583376f01aba9207c92e0d6a7d4e4f358a8964aa"},
{file = "haystack_experimental-0.2.0.tar.gz", hash = "sha256:393c543f29c50ea21365f5d032cd8b6bca10256016a3791546f1acc7ac623b07"},
]
[package.dependencies]
@@ -2457,18 +2415,18 @@ files = [
[[package]]
name = "langchain"
version = "0.3.0"
version = "0.3.1"
description = "Building applications with LLMs through composability"
optional = false
python-versions = "<4.0,>=3.9"
files = [
{file = "langchain-0.3.0-py3-none-any.whl", hash = "sha256:59a75a6a1eb7bfd2a6bf0c7a5816409a8fdc9046187b07af287b23b9899617af"},
{file = "langchain-0.3.0.tar.gz", hash = "sha256:a7c23892440bd1f5b9e029ff0dd709dd881ae927c4c0a3210ac64dba9bbf3f7f"},
{file = "langchain-0.3.1-py3-none-any.whl", hash = "sha256:94e5ee7464d4366e4b158aa5704953c39701ea237b9ed4b200096d49e83bb3ae"},
{file = "langchain-0.3.1.tar.gz", hash = "sha256:54d6e3abda2ec056875a231a418a4130ba7576e629e899067e499bfc847b7586"},
]
[package.dependencies]
aiohttp = ">=3.8.3,<4.0.0"
langchain-core = ">=0.3.0,<0.4.0"
langchain-core = ">=0.3.6,<0.4.0"
langchain-text-splitters = ">=0.3.0,<0.4.0"
langsmith = ">=0.1.17,<0.2.0"
numpy = {version = ">=1.26.0,<2.0.0", markers = "python_version >= \"3.12\""}
@@ -2480,21 +2438,21 @@ tenacity = ">=8.1.0,<8.4.0 || >8.4.0,<9.0.0"
[[package]]
name = "langchain-community"
version = "0.3.0"
version = "0.3.1"
description = "Community contributed LangChain integrations."
optional = false
python-versions = "<4.0,>=3.9"
files = [
{file = "langchain_community-0.3.0-py3-none-any.whl", hash = "sha256:40084f1f785f0fb56c8698ff059bbda8bd8c683cbdffa7902a0e04e72961496c"},
{file = "langchain_community-0.3.0.tar.gz", hash = "sha256:1ee8a469ad66977f21b9d96bdcdd8549c5281c474f0f9cc13b932efd63a78105"},
{file = "langchain_community-0.3.1-py3-none-any.whl", hash = "sha256:627eb26c16417764762ac47dd0d3005109f750f40242a88bb8f2958b798bcf90"},
{file = "langchain_community-0.3.1.tar.gz", hash = "sha256:c964a70628f266a61647e58f2f0434db633d4287a729f100a81dd8b0654aec93"},
]
[package.dependencies]
aiohttp = ">=3.8.3,<4.0.0"
dataclasses-json = ">=0.5.7,<0.7"
langchain = ">=0.3.0,<0.4.0"
langchain-core = ">=0.3.0,<0.4.0"
langsmith = ">=0.1.112,<0.2.0"
langchain = ">=0.3.1,<0.4.0"
langchain-core = ">=0.3.6,<0.4.0"
langsmith = ">=0.1.125,<0.2.0"
numpy = {version = ">=1.26.0,<2.0.0", markers = "python_version >= \"3.12\""}
pydantic-settings = ">=2.4.0,<3.0.0"
PyYAML = ">=5.3"
@@ -2504,13 +2462,13 @@ tenacity = ">=8.1.0,<8.4.0 || >8.4.0,<9.0.0"
[[package]]
name = "langchain-core"
version = "0.3.5"
version = "0.3.6"
description = "Building applications with LLMs through composability"
optional = false
python-versions = "<4.0,>=3.9"
files = [
{file = "langchain_core-0.3.5-py3-none-any.whl", hash = "sha256:2b5f86c1101beb013cb264c5722ad21931641493b4dc86e6f0575da698bf5cff"},
{file = "langchain_core-0.3.5.tar.gz", hash = "sha256:67e5510559454f3f7a0526e7ef91fd0f12b45c0cdc70720e44909f62b5becf5a"},
{file = "langchain_core-0.3.6-py3-none-any.whl", hash = "sha256:7bb3df0117bdc628b18b6c8748de72c6f537d745d47566053ce6650d5712281c"},
{file = "langchain_core-0.3.6.tar.gz", hash = "sha256:eb190494a5483f1965f693bb2085edb523370b20fc52dc294d3bd425773cd076"},
]
[package.dependencies]
@@ -3712,13 +3670,13 @@ virtualenv = ">=20.10.0"
[[package]]
name = "prompt-toolkit"
version = "3.0.47"
version = "3.0.48"
description = "Library for building powerful interactive command lines in Python"
optional = false
python-versions = ">=3.7.0"
files = [
{file = "prompt_toolkit-3.0.47-py3-none-any.whl", hash = "sha256:0d7bfa67001d5e39d02c224b663abc33687405033a8c422d0d675a5a13361d10"},
{file = "prompt_toolkit-3.0.47.tar.gz", hash = "sha256:1e1b29cb58080b1e69f207c893a1a7bf16d127a5c30c9d17a25a5d77792e5360"},
{file = "prompt_toolkit-3.0.48-py3-none-any.whl", hash = "sha256:f49a827f90062e411f1ce1f854f2aedb3c23353244f8108b89283587397ac10e"},
{file = "prompt_toolkit-3.0.48.tar.gz", hash = "sha256:d6623ab0477a80df74e646bdbc93621143f5caf104206aa29294d53de1a03d90"},
]
[package.dependencies]
@@ -6156,4 +6114,4 @@ testing = ["coverage (>=5.0.3)", "zope.event", "zope.testing"]
[metadata]
lock-version = "2.0"
python-versions = ">=3.12.*, <4.0"
content-hash = "368709143cec64b200c09e6ecf36a4bb50fbc4da963dfe58a6492dd74d6bcf83"
content-hash = "b0f4726eec24baccb3bd191828ca488a521fe1de31f5bf34067f165ff58b6f72"
+2
View File
@@ -28,6 +28,8 @@ ollama = "==0.2.1"
toml = "==0.10.2"
sqlglot = "==25.18.0"
cachetools = "==5.5.0"
pyyaml = "==6.0.2"
pydantic-settings = "==2.5.2"
[tool.poetry.group.dev.dependencies]
pre-commit = "==3.7.1"
+6 -9
View File
@@ -9,14 +9,13 @@ from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import ORJSONResponse, RedirectResponse
from langfuse.decorators import langfuse_context
from src.core.engine import EngineConfig
from src.globals import (
create_service_container,
create_service_metadata,
)
from src.providers import generate_components
from src.utils import (
init_langfuse,
init_providers,
load_env_vars,
setup_custom_logger,
)
@@ -35,12 +34,10 @@ setup_custom_logger(
@asynccontextmanager
async def lifespan(app: FastAPI):
# startup events
providers = init_providers(
engine_config=EngineConfig(provider=os.getenv("ENGINE", "wren_ui"))
)
pipe_components = generate_components()
app.state.service_container = create_service_container(
*providers,
should_force_deploy=bool(os.getenv("SHOULD_FORCE_DEPLOY", "")),
pipe_components,
column_indexing_batch_size=(
int(os.getenv("COLUMN_INDEXING_BATCH_SIZE"))
if os.getenv("COLUMN_INDEXING_BATCH_SIZE")
@@ -63,7 +60,7 @@ async def lifespan(app: FastAPI):
"ttl": int(os.getenv("QUERY_CACHE_TTL") or 120),
},
)
app.state.service_metadata = create_service_metadata(*providers)
app.state.service_metadata = create_service_metadata(pipe_components)
init_langfuse()
yield
@@ -134,7 +131,7 @@ if __name__ == "__main__":
host=server_host,
port=server_port,
reload=should_reload,
reload_includes=["src/**/*.py", ".env.dev"],
reload_includes=["src/**/*.py", ".env.dev", "config.yaml"],
workers=1,
loop="uvloop",
http="httptools",
+22
View File
@@ -1,10 +1,15 @@
import asyncio
from abc import ABCMeta, abstractmethod
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, Dict
from hamilton.experimental.h_async import AsyncDriver
from haystack import Pipeline
from src.core.engine import Engine
from src.core.provider import DocumentStoreProvider, EmbedderProvider, LLMProvider
class BasicPipeline(metaclass=ABCMeta):
def __init__(self, pipe: Pipeline | AsyncDriver):
@@ -19,3 +24,20 @@ def async_validate(task: callable):
result = asyncio.run(task())
print(result)
return result
@dataclass
class PipelineComponent(Mapping):
llm_provider: LLMProvider = None
embedder_provider: EmbedderProvider = None
document_store_provider: DocumentStoreProvider = None
engine: Engine = None
def __getitem__(self, key):
return getattr(self, key)
def __iter__(self):
return iter(self.__dict__)
def __len__(self):
return len(self.__dict__)
+50 -56
View File
@@ -1,11 +1,11 @@
import logging
from dataclasses import dataclass
from dataclasses import asdict, dataclass
from typing import Optional
import toml
from src.core.engine import Engine
from src.core.provider import DocumentStoreProvider, EmbedderProvider, LLMProvider
from src.core.pipeline import PipelineComponent
from src.core.provider import EmbedderProvider, LLMProvider
from src.pipelines.generation import (
followup_sql_generation,
sql_answer,
@@ -43,36 +43,22 @@ class ServiceContainer:
@dataclass
class ServiceMetadata:
models_metadata: dict
pipes_metadata: dict
service_version: str
def create_service_container(
llm_provider: LLMProvider,
embedder_provider: EmbedderProvider,
document_store_provider: DocumentStoreProvider,
engine: Engine,
should_force_deploy: Optional[str] = None,
pipe_components: dict[str, PipelineComponent],
column_indexing_batch_size: Optional[int] = 50,
table_retrieval_size: Optional[int] = 10,
table_column_retrieval_size: Optional[int] = 1000,
query_cache: Optional[dict] = {},
) -> ServiceContainer:
if should_force_deploy:
document_store_provider.get_store(recreate_index=True)
document_store_provider.get_store(
dataset_name="table_descriptions", recreate_index=True
)
document_store_provider.get_store(
dataset_name="view_questions", recreate_index=True
)
return ServiceContainer(
semantics_preparation_service=SemanticsPreparationService(
pipelines={
"indexing": indexing.Indexing(
embedder_provider=embedder_provider,
document_store_provider=document_store_provider,
**pipe_components["indexing"],
column_indexing_batch_size=column_indexing_batch_size,
),
},
@@ -81,30 +67,24 @@ def create_service_container(
ask_service=AskService(
pipelines={
"retrieval": retrieval.Retrieval(
llm_provider=llm_provider,
embedder_provider=embedder_provider,
document_store_provider=document_store_provider,
**pipe_components["retrieval"],
table_retrieval_size=table_retrieval_size,
table_column_retrieval_size=table_column_retrieval_size,
),
"historical_question": historical_question.HistoricalQuestion(
embedder_provider=embedder_provider,
store_provider=document_store_provider,
**pipe_components["historical_question"],
),
"sql_generation": sql_generation.SQLGeneration(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_generation"],
),
"sql_correction": sql_correction.SQLCorrection(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_correction"],
),
"followup_sql_generation": followup_sql_generation.FollowUpSQLGeneration(
llm_provider=llm_provider,
engine=engine,
**pipe_components["followup_sql_generation"],
),
"sql_summary": sql_summary.SQLSummary(
llm_provider=llm_provider,
**pipe_components["sql_summary"],
),
},
**query_cache,
@@ -112,8 +92,7 @@ def create_service_container(
sql_answer_service=SqlAnswerService(
pipelines={
"sql_answer": sql_answer.SQLAnswer(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_answer"],
)
},
**query_cache,
@@ -121,8 +100,7 @@ def create_service_container(
ask_details_service=AskDetailsService(
pipelines={
"sql_breakdown": sql_breakdown.SQLBreakdown(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_breakdown"],
),
},
**query_cache,
@@ -130,22 +108,18 @@ def create_service_container(
sql_expansion_service=SqlExpansionService(
pipelines={
"retrieval": retrieval.Retrieval(
llm_provider=llm_provider,
embedder_provider=embedder_provider,
document_store_provider=document_store_provider,
**pipe_components["retrieval"],
table_retrieval_size=table_retrieval_size,
table_column_retrieval_size=table_column_retrieval_size,
),
"sql_expansion": sql_expansion.SQLExpansion(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_expansion"],
),
"sql_correction": sql_correction.SQLCorrection(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_correction"],
),
"sql_summary": sql_summary.SQLSummary(
llm_provider=llm_provider,
**pipe_components["sql_summary"],
),
},
**query_cache,
@@ -153,7 +127,7 @@ def create_service_container(
sql_explanation_service=SQLExplanationService(
pipelines={
"sql_explanation": sql_explanation.SQLExplanation(
llm_provider=llm_provider,
**pipe_components["sql_explanation"],
)
},
**query_cache,
@@ -161,8 +135,7 @@ def create_service_container(
sql_regeneration_service=SQLRegenerationService(
pipelines={
"sql_regeneration": sql_regeneration.SQLRegeneration(
llm_provider=llm_provider,
engine=engine,
**pipe_components["sql_regeneration"],
)
},
**query_cache,
@@ -178,9 +151,7 @@ def get_service_container():
def create_service_metadata(
llm_provider: LLMProvider,
embedder_provider: EmbedderProvider,
*_,
pipe_components: dict[str, PipelineComponent],
pyproject_path: str = "pyproject.toml",
) -> ServiceMetadata:
def _get_version_from_pyproject() -> str:
@@ -188,17 +159,40 @@ def create_service_metadata(
pyproject = toml.load(f)
return pyproject["tool"]["poetry"]["version"]
models_metadata = {
"generation_model": llm_provider.get_model(),
"generation_model_kwargs": llm_provider.get_model_kwargs(),
"embedding_model": embedder_provider.get_model(),
"embedding_model_dim": embedder_provider.get_dimensions(),
def _convert_pipe_metadata(
llm_provider: LLMProvider,
embedder_provider: EmbedderProvider,
**_,
) -> dict:
llm_metadata = (
{
"llm_model": llm_provider.get_model(),
"llm_model_kwargs": llm_provider.get_model_kwargs(),
}
if llm_provider
else {}
)
embedding_metadata = (
{
"embedding_model": embedder_provider.get_model(),
"embedding_model_dim": embedder_provider.get_dimensions(),
}
if embedder_provider
else {}
)
return {**llm_metadata, **embedding_metadata}
pipes_metadata = {
pipe_name: _convert_pipe_metadata(**asdict(component))
for pipe_name, component in pipe_components.items()
}
service_version = _get_version_from_pyproject()
logger.info(f"Service version: {service_version}")
return ServiceMetadata(models_metadata, service_version)
return ServiceMetadata(pipes_metadata, service_version)
# Create a dependency that will be used to access the ServiceMetadata
@@ -179,6 +179,7 @@ class FollowUpSQLGeneration(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"generator": llm_provider.get_generator(
@@ -256,7 +257,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -164,6 +164,7 @@ class SQLAnswer(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"data_fetcher": DataFetcher(engine=engine),
@@ -224,7 +225,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -149,6 +149,7 @@ class SQLBreakdown(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"generator": llm_provider.get_generator(
@@ -202,7 +203,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -110,6 +110,7 @@ class SQLCorrection(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"generator": llm_provider.get_generator(
@@ -179,7 +180,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -94,6 +94,7 @@ class SQLExpansion(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"generator": llm_provider.get_generator(
@@ -161,7 +162,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -590,6 +590,7 @@ class SQLExplanation(BasicPipeline):
def __init__(
self,
llm_provider: LLMProvider,
**kwargs,
):
self._components = {
"pre_processor": SQLAnalysisPreprocessor(),
@@ -655,7 +656,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -147,6 +147,7 @@ class SQLGeneration(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"generator": llm_provider.get_generator(
@@ -224,7 +225,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -158,6 +158,7 @@ class SQLRegeneration(BasicPipeline):
self,
llm_provider: LLMProvider,
engine: Engine,
**kwargs,
):
self._components = {
"preprocesser": SQLRegenerationPreprocesser(),
@@ -222,7 +223,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -121,6 +121,7 @@ class SQLSummary(BasicPipeline):
def __init__(
self,
llm_provider: LLMProvider,
**kwargs,
):
self._components = {
"generator": llm_provider.get_generator(
@@ -178,7 +179,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -625,6 +625,7 @@ class Indexing(BasicPipeline):
embedder_provider: EmbedderProvider,
document_store_provider: DocumentStoreProvider,
column_indexing_batch_size: Optional[int] = 50,
**kwargs,
) -> None:
dbschema_store = document_store_provider.get_store()
view_store = document_store_provider.get_store(dataset_name="view_questions")
@@ -701,7 +702,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -141,13 +141,14 @@ class HistoricalQuestion(BasicPipeline):
def __init__(
self,
embedder_provider: EmbedderProvider,
store_provider: DocumentStoreProvider,
document_store_provider: DocumentStoreProvider,
**kwargs,
) -> None:
store = store_provider.get_store(dataset_name="view_questions")
store = document_store_provider.get_store(dataset_name="view_questions")
self._components = {
"store": store,
"embedder": embedder_provider.get_text_embedder(),
"retriever": store_provider.get_retriever(
"retriever": document_store_provider.get_retriever(
document_store=store,
),
"score_filter": ScoreFilter(),
@@ -199,7 +200,8 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_providers
from src.utils import init_langfuse, load_env_vars
load_env_vars()
init_langfuse()
@@ -209,7 +211,8 @@ if __name__ == "__main__":
)
pipeline = HistoricalQuestion(
embedder_provider=embedder_provider, store_provider=document_store_provider
embedder_provider=embedder_provider,
document_store_provider=document_store_provider,
)
pipeline.visualize("this is a query")
@@ -292,6 +292,7 @@ class Retrieval(BasicPipeline):
document_store_provider: DocumentStoreProvider,
table_retrieval_size: Optional[int] = 10,
table_column_retrieval_size: Optional[int] = 1000,
**kwargs,
):
self._components = {
"embedder": embedder_provider.get_text_embedder(),
@@ -355,7 +356,7 @@ if __name__ == "__main__":
from src.core.engine import EngineConfig
from src.core.pipeline import async_validate
from src.utils import init_langfuse, init_providers, load_env_vars
from src.providers import init_langfuse, init_providers, load_env_vars
load_env_vars()
init_langfuse()
+197
View File
@@ -0,0 +1,197 @@
import logging
import os
from collections.abc import Mapping
from typing import Tuple
import yaml
from yaml.loader import SafeLoader
from src.core.engine import Engine, EngineConfig
from src.core.pipeline import PipelineComponent
from src.core.provider import DocumentStoreProvider, EmbedderProvider, LLMProvider
from src.providers import loader
logger = logging.getLogger("wren-ai-service")
def provider_factory(
config: dict = {},
) -> LLMProvider | EmbedderProvider | DocumentStoreProvider | Engine:
return loader.get_provider(config.get("provider"))(**config)
def load_config(path: str = "config.yaml") -> list[dict]:
if not os.path.exists(path):
return []
with open(path, "r") as f:
return list(yaml.load_all(f, Loader=SafeLoader))
def process_llm(entry: dict) -> dict:
others = {
k: v
for k, v in entry.items()
if k not in ["type", "provider", "api_key", "models"]
}
returned = {}
for model in entry["models"]:
model_name = f"{entry['provider']}.{model['model']}"
returned[model_name] = {
"provider": entry["provider"],
"api_key": entry["api_key"],
"model": model["model"],
"kwargs": model["kwargs"],
**others,
}
return returned
def process_embedder(entry: dict) -> dict:
others = {
k: v
for k, v in entry.items()
if k not in ["type", "provider", "api_key", "models"]
}
returned = {}
for model in entry["models"]:
model_name = f"{entry['provider']}.{model['model']}"
returned[model_name] = {
"provider": entry["provider"],
"api_key": entry["api_key"],
"model": model["model"],
"dimension": model["dimension"],
**others,
}
return returned
def process_document_store(entry: dict) -> dict:
return {entry["provider"]: {k: v for k, v in entry.items() if k not in ["type"]}}
def process_engine(entry: dict) -> dict:
return {entry["provider"]: {k: v for k, v in entry.items() if k not in ["type"]}}
def process_pipeline(entry: dict) -> dict:
return {
pipe["name"]: {
"llm": pipe.get("llm"),
"embedder": pipe.get("embedder"),
"document_store": pipe.get("document_store"),
"engine": pipe.get("engine"),
}
for pipe in entry["pipes"]
}
def convert_data(config: list[dict]) -> dict:
returned = {
"embedder": {},
"llm": {},
"document_store": {},
"engine": {},
"pipeline": {},
}
type_to_processor = {
"llm": process_llm,
"embedder": process_embedder,
"document_store": process_document_store,
"engine": process_engine,
"pipeline": process_pipeline,
}
for entry in config:
type = entry["type"]
processor = type_to_processor.get(type)
if not processor:
logger.error(f"Unknown type: {type}")
raise ValueError(f"Unknown type: {type}")
converted = processor(entry)
returned[type].update(converted)
return returned
def init_providers(
engine_config: EngineConfig,
) -> Tuple[LLMProvider, EmbedderProvider, DocumentStoreProvider, Engine]:
logger.info("Initializing providers...")
loader.import_mods()
llm_provider = loader.get_provider(os.getenv("LLM_PROVIDER", "openai_llm"))()
embedder_provider = loader.get_provider(
os.getenv("EMBEDDER_PROVIDER", "openai_embedder")
)()
document_store_provider = loader.get_provider(
os.getenv("DOCUMENT_STORE_PROVIDER", "qdrant")
)()
engine = loader.get_provider(engine_config.provider)(**engine_config.config)
return llm_provider, embedder_provider, document_store_provider, engine
class Wrapper(Mapping):
def __init__(self):
self.value = PipelineComponent(
*init_providers(
engine_config=EngineConfig(provider=os.getenv("ENGINE", "wren_ui"))
)
)
def __getitem__(self, key):
return self.value
def __repr__(self):
return f"Wrapper({self.value})"
def __iter__(self):
return iter(self.value)
def __len__(self):
return len(self.value)
def generate_components() -> dict[str, PipelineComponent]:
raw = load_config()
if not raw:
# if no config, initialize the providers from the environment variables
return Wrapper()
config = convert_data(raw)
loader.import_mods()
providers = {
"embedder": config.get("embedder", {}),
"llm": config.get("llm", {}),
"document_store": config.get("document_store", {}),
"engine": config.get("engine", {}),
}
instantiated_providers = {
category: {
identifier: provider_factory(config)
for identifier, config in configs.items()
}
for category, configs in providers.items()
}
def get(type: str, components: dict):
return instantiated_providers[type].get(components.get(type))
def componentize(components: dict):
return PipelineComponent(
embedder_provider=get("embedder", components),
llm_provider=get("llm", components),
document_store_provider=get("document_store", components),
engine=get("engine", components),
)
return {
pipe_name: componentize(components)
for pipe_name, components in config.get("pipeline", {}).items()
}
@@ -333,19 +333,10 @@ class QdrantProvider(DocumentStoreProvider):
def __init__(
self,
location: str = os.getenv("QDRANT_HOST", "qdrant"),
api_key: Optional[Secret] = Secret.from_env_var("QDRANT_API_KEY")
if os.getenv("QDRANT_API_KEY")
else None,
api_key: Optional[str] = os.getenv("QDRANT_API_KEY", None),
timeout: Optional[int] = (
int(os.getenv("QDRANT_TIMEOUT")) if os.getenv("QDRANT_TIMEOUT") else 120
),
):
self._location = location
self._api_key = api_key
self._timeout = timeout
def get_store(
self,
embedding_model_dim: int = (
int(os.getenv("EMBEDDING_MODEL_DIMENSION"))
if os.getenv("EMBEDDING_MODEL_DIMENSION")
@@ -354,17 +345,26 @@ class QdrantProvider(DocumentStoreProvider):
or get_default_embedding_model_dim(
os.getenv("EMBEDDER_PROVIDER", "openai_embedder")
),
**_,
):
self._location = location
self._api_key = Secret.from_token(api_key) if api_key else None
self._timeout = timeout
self._embedding_model_dim = embedding_model_dim
def get_store(
self,
dataset_name: Optional[str] = None,
recreate_index: bool = False,
):
logger.info(
f"Using Qdrant Document Store with Embedding Model Dimension: {embedding_model_dim}"
f"Using Qdrant Document Store with Embedding Model Dimension: {self._embedding_model_dim}"
)
return AsyncQdrantDocumentStore(
location=self._location,
api_key=self._api_key,
embedding_dim=embedding_model_dim,
embedding_dim=self._embedding_model_dim,
index=dataset_name or "Document",
recreate_index=recreate_index,
on_disk=True,
@@ -375,7 +375,7 @@ class QdrantProvider(DocumentStoreProvider):
always_ram=True,
)
)
if embedding_model_dim >= 1024
if self._embedding_model_dim >= 1024
else None
),
# to improve the indexing performance, we disable building global index for the whole collection
@@ -195,11 +195,11 @@ class AsyncDocumentEmbedder(AzureOpenAIDocumentEmbedder):
class AzureOpenAIEmbedderProvider(EmbedderProvider):
def __init__(
self,
embed_api_key: Secret = Secret.from_env_var("EMBEDDER_AZURE_OPENAI_API_KEY"),
embed_api_base: str = os.getenv("EMBEDDER_AZURE_OPENAI_API_BASE"),
embed_api_version: str = os.getenv("EMBEDDER_AZURE_OPENAI_VERSION"),
embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL,
embedding_model_dim: int = (
api_key: Secret = Secret.from_env_var("EMBEDDER_AZURE_OPENAI_API_KEY"),
api_base: str = os.getenv("EMBEDDER_AZURE_OPENAI_API_BASE"),
api_version: str = os.getenv("EMBEDDER_AZURE_OPENAI_VERSION"),
model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL,
dimension: int = (
int(os.getenv("EMBEDDING_MODEL_DIMENSION"))
if os.getenv("EMBEDDING_MODEL_DIMENSION")
else 0
@@ -210,12 +210,13 @@ class AzureOpenAIEmbedderProvider(EmbedderProvider):
if os.getenv("EMBEDDER_TIMEOUT")
else 120.0
),
**_,
):
self._embedding_api_base = remove_trailing_slash(embed_api_base)
self._embedding_api_key = embed_api_key
self._embedding_api_version = embed_api_version
self._embedding_model = embedding_model
self._embedding_model_dim = embedding_model_dim
self._embedding_api_base = remove_trailing_slash(api_base)
self._embedding_api_key = api_key
self._embedding_api_version = api_version
self._embedding_model = model
self._embedding_model_dim = dimension
self._timeout = timeout
logger.info(f"Using Azure OpenAI Embedding Model: {self._embedding_model}")
@@ -162,8 +162,8 @@ class OllamaEmbedderProvider(EmbedderProvider):
def __init__(
self,
url: str = os.getenv("EMBEDDER_OLLAMA_URL") or EMBEDDER_OLLAMA_URL,
embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL,
embedding_model_dim: int = (
model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL,
dimension: int = (
int(os.getenv("EMBEDDING_MODEL_DIMENSION"))
if os.getenv("EMBEDDING_MODEL_DIMENSION")
else 0
@@ -172,10 +172,11 @@ class OllamaEmbedderProvider(EmbedderProvider):
timeout: Optional[int] = (
int(os.getenv("EMBEDDER_TIMEOUT")) if os.getenv("EMBEDDER_TIMEOUT") else 120
),
**_,
):
self._url = remove_trailing_slash(url)
self._embedding_model = embedding_model
self._embedding_model_dim = embedding_model_dim
self._embedding_model = model
self._embedding_model_dim = dimension
self._timeout = timeout
pull_ollama_model(self._url, self._embedding_model)
@@ -183,11 +183,11 @@ class AsyncDocumentEmbedder(OpenAIDocumentEmbedder):
class OpenAIEmbedderProvider(EmbedderProvider):
def __init__(
self,
api_key: Secret = Secret.from_env_var("EMBEDDER_OPENAI_API_KEY"),
api_key: str = os.getenv("EMBEDDER_OPENAI_API_KEY"),
api_base: str = os.getenv("EMBEDDER_OPENAI_API_BASE")
or EMBEDDER_OPENAI_API_BASE,
embedding_model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL,
embedding_model_dim: int = (
model: str = os.getenv("EMBEDDING_MODEL") or EMBEDDING_MODEL,
dimension: int = (
int(os.getenv("EMBEDDING_MODEL_DIMENSION"))
if os.getenv("EMBEDDING_MODEL_DIMENSION")
else 0
@@ -198,11 +198,12 @@ class OpenAIEmbedderProvider(EmbedderProvider):
if os.getenv("EMBEDDER_TIMEOUT")
else 120.0
),
**_,
):
self._api_key = api_key
self._api_key = Secret.from_token(api_key)
self._api_base = remove_trailing_slash(api_base)
self._embedding_model = embedding_model
self._embedding_model_dim = embedding_model_dim
self._embedding_model = model
self._embedding_model_dim = dimension
self._timeout = timeout
logger.info(
+15 -8
View File
@@ -14,7 +14,11 @@ logger = logging.getLogger("wren-ai-service")
@provider("wren_ui")
class WrenUI(Engine):
def __init__(self, endpoint: str = os.getenv("WREN_UI_ENDPOINT")):
def __init__(
self,
endpoint: str = os.getenv("WREN_UI_ENDPOINT"),
**_,
):
self._endpoint = endpoint
logger.info("Using Engine: wren_ui")
@@ -60,16 +64,15 @@ class WrenIbis(Engine):
endpoint: str = os.getenv("WREN_IBIS_ENDPOINT"),
source: str = os.getenv("WREN_IBIS_SOURCE"),
manifest: str = os.getenv("WREN_IBIS_MANIFEST"),
connection_info: dict = (
orjson.loads(base64.b64decode(os.getenv("WREN_IBIS_CONNECTION_INFO")))
if os.getenv("WREN_IBIS_CONNECTION_INFO")
else {}
),
connection_info: str = os.getenv("WREN_IBIS_CONNECTION_INFO"),
**_,
):
self._endpoint = endpoint
self._source = source
self._manifest = manifest
self._connection_info = connection_info
self._connection_info = (
orjson.loads(base64.b64decode(connection_info)) if connection_info else {}
)
logger.info("Using Engine: wren_ibis")
async def execute_sql(
@@ -108,7 +111,11 @@ class WrenIbis(Engine):
@provider("wren_engine")
class WrenEngine(Engine):
def __init__(self, endpoint: str = os.getenv("WREN_ENGINE_ENDPOINT")):
def __init__(
self,
endpoint: str = os.getenv("WREN_ENGINE_ENDPOINT"),
**_,
):
self._endpoint = endpoint
logger.info("Using Engine: wren_engine")
@@ -120,11 +120,11 @@ class AsyncGenerator(AzureOpenAIGenerator):
class AzureOpenAILLMProvider(LLMProvider):
def __init__(
self,
chat_api_key: Secret = Secret.from_env_var("LLM_AZURE_OPENAI_API_KEY"),
chat_api_base: str = os.getenv("LLM_AZURE_OPENAI_API_BASE"),
chat_api_version: str = os.getenv("LLM_AZURE_OPENAI_VERSION"),
generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL,
model_kwargs: Dict[str, Any] = (
api_key: Secret = Secret.from_env_var("LLM_AZURE_OPENAI_API_KEY"),
api_base: str = os.getenv("LLM_AZURE_OPENAI_API_BASE"),
api_version: str = os.getenv("LLM_AZURE_OPENAI_VERSION"),
model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL,
kwargs: Dict[str, Any] = (
orjson.loads(os.getenv("GENERATION_MODEL_KWARGS"))
if os.getenv("GENERATION_MODEL_KWARGS")
else GENERATION_MODEL_KWARGS
@@ -132,12 +132,13 @@ class AzureOpenAILLMProvider(LLMProvider):
timeout: Optional[float] = (
float(os.getenv("LLM_TIMEOUT")) if os.getenv("LLM_TIMEOUT") else 120.0
),
**_,
):
self._generation_api_key = chat_api_key
self._generation_api_base = remove_trailing_slash(chat_api_base)
self._generation_api_version = chat_api_version
self._generation_model = generation_model
self._model_kwargs = model_kwargs
self._generation_api_key = api_key
self._generation_api_base = remove_trailing_slash(api_base)
self._generation_api_version = api_version
self._generation_model = model
self._model_kwargs = kwargs
self._timeout = timeout
logger.info(f"Using AzureOpenAI LLM: {self._generation_model}")
+5 -4
View File
@@ -125,8 +125,8 @@ class OllamaLLMProvider(LLMProvider):
def __init__(
self,
url: str = os.getenv("LLM_OLLAMA_URL") or LLM_OLLAMA_URL,
generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL,
model_kwargs: Dict[str, Any] = (
model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL,
kwargs: Dict[str, Any] = (
orjson.loads(os.getenv("GENERATION_MODEL_KWARGS"))
if os.getenv("GENERATION_MODEL_KWARGS")
else GENERATION_MODEL_KWARGS
@@ -134,10 +134,11 @@ class OllamaLLMProvider(LLMProvider):
timeout: int = (
int(os.getenv("LLM_TIMEOUT")) if os.getenv("LLM_TIMEOUT") else 120
),
**_,
):
self._url = remove_trailing_slash(url)
self._generation_model = generation_model
self._model_kwargs = model_kwargs
self._generation_model = model
self._model_kwargs = kwargs
self._timeout = timeout
pull_ollama_model(self._url, self._generation_model)
+7 -6
View File
@@ -120,10 +120,10 @@ class AsyncGenerator(OpenAIGenerator):
class OpenAILLMProvider(LLMProvider):
def __init__(
self,
api_key: Secret = Secret.from_env_var("LLM_OPENAI_API_KEY"),
api_key: str = os.getenv("LLM_OPENAI_API_KEY"),
api_base: str = os.getenv("LLM_OPENAI_API_BASE") or LLM_OPENAI_API_BASE,
generation_model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL,
model_kwargs: Dict[str, Any] = (
model: str = os.getenv("GENERATION_MODEL") or GENERATION_MODEL,
kwargs: Dict[str, Any] = (
orjson.loads(os.getenv("GENERATION_MODEL_KWARGS"))
if os.getenv("GENERATION_MODEL_KWARGS")
else GENERATION_MODEL_KWARGS
@@ -131,11 +131,12 @@ class OpenAILLMProvider(LLMProvider):
timeout: Optional[float] = (
float(os.getenv("LLM_TIMEOUT")) if os.getenv("LLM_TIMEOUT") else 120.0
),
**_,
):
self._api_key = api_key
self._api_key = Secret.from_token(api_key)
self._api_base = remove_trailing_slash(api_base)
self._generation_model = generation_model
self._model_kwargs = model_kwargs
self._generation_model = model
self._model_kwargs = kwargs
self._timeout = timeout
logger.info(f"Using OpenAILLM provider with API base: {self._api_base}")
+2 -25
View File
@@ -4,15 +4,10 @@ import logging
import os
import time
from pathlib import Path
from typing import Tuple
from dotenv import load_dotenv
from langfuse.decorators import langfuse_context
from src.core.engine import Engine, EngineConfig
from src.core.provider import DocumentStoreProvider, EmbedderProvider, LLMProvider
from src.providers import loader
logger = logging.getLogger("wren-ai-service")
@@ -58,24 +53,6 @@ def load_env_vars() -> str:
return "prod"
def init_providers(
engine_config: EngineConfig,
) -> Tuple[LLMProvider, EmbedderProvider, DocumentStoreProvider, Engine]:
logger.info("Initializing providers...")
loader.import_mods()
llm_provider = loader.get_provider(os.getenv("LLM_PROVIDER", "openai_llm"))()
embedder_provider = loader.get_provider(
os.getenv("EMBEDDER_PROVIDER", "openai_embedder")
)()
document_store_provider = loader.get_provider(
os.getenv("DOCUMENT_STORE_PROVIDER", "qdrant")
)()
engine = loader.get_provider(engine_config.provider)(**engine_config.config)
return llm_provider, embedder_provider, document_store_provider, engine
def timer(func):
@functools.wraps(func)
def wrapper_timer(*args, **kwargs):
@@ -186,12 +163,12 @@ def trace_metadata(func):
service_metadata = kwargs.get(
"service_metadata",
{
"models_metadata": {},
"pipes_metadata": {},
"service_version": "",
},
)
langfuse_metadata = {
**service_metadata.get("models_metadata"),
**service_metadata.get("pipes_metadata"),
**addition,
"mdl_hash": metadata.get("mdl_hash"),
"project_id": metadata.get("project_id"),
@@ -10,7 +10,7 @@ from src.pipelines.generation.sql_correction import SQLCorrection
from src.pipelines.generation.sql_generation import SQLGeneration
from src.pipelines.indexing.indexing import Indexing
from src.pipelines.retrieval.retrieval import Retrieval
from src.utils import init_providers
from src.providers import init_providers
from src.web.v1.services.ask import AskHistory
from src.web.v1.services.ask_details import SQLBreakdown
@@ -2,7 +2,7 @@ import pytest
from src.core.engine import EngineConfig
from src.pipelines.generation import sql_breakdown
from src.utils import init_providers
from src.providers import init_providers
@pytest.mark.asyncio
@@ -4,22 +4,21 @@ from haystack.document_stores.types import DocumentStore
from src.core.engine import EngineConfig
from src.pipelines.indexing.indexing import DocumentCleaner
from src.utils import init_providers
from src.providers import init_providers
@pytest.mark.asyncio
async def _mock_store(name: str = "default") -> DocumentStore:
_, _, document_store_provider, _ = init_providers(EngineConfig())
store = document_store_provider.get_store(
embedding_model_dim=5,
dataset_name=name,
recreate_index=True,
)
await store.write_documents(
[
Document(id=str(0), content="This is first", embedding=[0.0] * 5),
Document(id=str(1), content="This is second", embedding=[0.1] * 5),
Document(id=str(0), content="This is first", embedding=[0.0] * 3072),
Document(id=str(1), content="This is second", embedding=[0.1] * 3072),
]
)
assert (await store.count_documents()) == 2
@@ -10,7 +10,7 @@ from src.core.engine import EngineConfig
from src.pipelines.generation import sql_correction, sql_generation
from src.pipelines.indexing import indexing
from src.pipelines.retrieval import historical_question, retrieval
from src.utils import init_providers
from src.providers import init_providers
from src.web.v1.services.ask import (
AskRequest,
AskResultRequest,
@@ -44,7 +44,7 @@ def ask_service():
),
"historical_question": historical_question.HistoricalQuestion(
embedder_provider=embedder_provider,
store_provider=document_store_provider,
document_store_provider=document_store_provider,
),
"sql_generation": sql_generation.SQLGeneration(
llm_provider=llm_provider,
@@ -75,11 +75,13 @@ def indexing_service():
@pytest.fixture
def service_metadata():
return {
"models_metadata": {
"generation_model": "mock-llm-model",
"generation_model_kwargs": {},
"embedding_model": "mock-embedding-model",
"embedding_model_dim": 768,
"pipes_metadata": {
"mock": {
"generation_model": "mock-llm-model",
"generation_model_kwargs": {},
"embedding_model": "mock-embedding-model",
"embedding_model_dim": 768,
},
},
"service_version": "0.8.0-mock",
}
@@ -4,7 +4,7 @@ import pytest
from src.core.engine import EngineConfig
from src.pipelines.generation import sql_breakdown
from src.utils import init_providers
from src.providers import init_providers
from src.web.v1.services.ask_details import (
AskDetailsRequest,
AskDetailsResultRequest,
@@ -28,11 +28,13 @@ def ask_details_service():
@pytest.fixture
def service_metadata():
return {
"models_metadata": {
"generation_model": "mock-llm-model",
"generation_model_kwargs": {},
"embedding_model": "mock-embedding-model",
"embedding_model_dim": 768,
"pipes_metadata": {
"mock": {
"generation_model": "mock-llm-model",
"generation_model_kwargs": {},
"embedding_model": "mock-embedding-model",
"embedding_model_dim": 768,
},
},
"service_version": "0.8.0-mock",
}
+10 -7
View File
@@ -6,6 +6,7 @@ import pytest
from pytest_mock import MockFixture
import src.utils as utils
from src.core.pipeline import PipelineComponent
from src.globals import ServiceMetadata, create_service_metadata
@@ -37,17 +38,19 @@ def service_metadata(mocker: MockFixture):
current_path = os.path.dirname(__file__)
return create_service_metadata(
*_mock(mocker),
pipe_components={"mock": PipelineComponent(*_mock(mocker))},
pyproject_path=os.path.join(current_path, "../data/mock_pyproject.toml"),
)
def test_service_metadata(service_metadata: ServiceMetadata):
assert service_metadata.models_metadata == {
"generation_model": "mock-llm-model",
"generation_model_kwargs": {},
"embedding_model": "mock-embedding-model",
"embedding_model_dim": 768,
assert service_metadata.pipes_metadata == {
"mock": {
"llm_model": "mock-llm-model",
"llm_model_kwargs": {},
"embedding_model": "mock-embedding-model",
"embedding_model_dim": 768,
},
}
assert service_metadata.service_version == "0.8.0-mock"
@@ -77,6 +80,6 @@ def test_trace_metadata(service_metadata: ServiceMetadata, mocker: MockFixture):
metadata={
"mdl_hash": "mock-mdl-hash",
"project_id": "mock-project-id",
**service_metadata.models_metadata,
**service_metadata.pipes_metadata,
},
)