From 7f4a0bebc7f8c6dc40d95ff41c2d31f025278880 Mon Sep 17 00:00:00 2001 From: Pao Sheng Date: Tue, 1 Oct 2024 14:08:24 +0800 Subject: [PATCH] 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 --- .gitignore | 1 + wren-ai-service/config.example.yaml | 105 ++++++++++ wren-ai-service/demo/app.py | 26 --- wren-ai-service/demo/utils.py | 120 ++++++----- wren-ai-service/eval/prediction.py | 3 +- wren-ai-service/poetry.lock | 82 ++------ wren-ai-service/pyproject.toml | 2 + wren-ai-service/src/__main__.py | 15 +- wren-ai-service/src/core/pipeline.py | 22 ++ wren-ai-service/src/globals.py | 106 +++++----- .../generation/followup_sql_generation.py | 4 +- .../src/pipelines/generation/sql_answer.py | 4 +- .../src/pipelines/generation/sql_breakdown.py | 4 +- .../pipelines/generation/sql_correction.py | 4 +- .../src/pipelines/generation/sql_expansion.py | 4 +- .../pipelines/generation/sql_explanation.py | 4 +- .../pipelines/generation/sql_generation.py | 4 +- .../pipelines/generation/sql_regeneration.py | 4 +- .../src/pipelines/generation/sql_summary.py | 4 +- .../src/pipelines/indexing/indexing.py | 4 +- .../retrieval/historical_question.py | 13 +- .../src/pipelines/retrieval/retrieval.py | 3 +- wren-ai-service/src/providers/__init__.py | 197 ++++++++++++++++++ .../src/providers/document_store/qdrant.py | 26 +-- .../src/providers/embedder/azure_openai.py | 21 +- .../src/providers/embedder/ollama.py | 9 +- .../src/providers/embedder/openai.py | 13 +- wren-ai-service/src/providers/engine/wren.py | 23 +- .../src/providers/llm/azure_openai.py | 21 +- wren-ai-service/src/providers/llm/ollama.py | 9 +- wren-ai-service/src/providers/llm/openai.py | 13 +- wren-ai-service/src/utils.py | 27 +-- .../tests/pytest/pipelines/test_ask.py | 2 +- .../pytest/pipelines/test_ask_details.py | 2 +- .../pytest/pipelines/test_document_cleaner.py | 7 +- .../tests/pytest/services/test_ask.py | 16 +- .../tests/pytest/services/test_ask_details.py | 14 +- wren-ai-service/tests/pytest/test_utils.py | 17 +- 38 files changed, 612 insertions(+), 343 deletions(-) create mode 100644 wren-ai-service/config.example.yaml diff --git a/.gitignore b/.gitignore index c30e36541..160a0a600 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/wren-ai-service/config.example.yaml b/wren-ai-service/config.example.yaml new file mode 100644 index 000000000..64c41add3 --- /dev/null +++ b/wren-ai-service/config.example.yaml @@ -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 diff --git a/wren-ai-service/demo/app.py b/wren-ai-service/demo/app.py index cb353e8c4..2f5a60d4a 100644 --- a/wren-ai-service/demo/app.py +++ b/wren-ai-service/demo/app.py @@ -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( diff --git a/wren-ai-service/demo/utils.py b/wren-ai-service/demo/utils.py index f36eebf54..fd71726bf 100644 --- a/wren-ai-service/demo/utils.py +++ b/wren-ai-service/demo/utils.py @@ -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 diff --git a/wren-ai-service/eval/prediction.py b/wren-ai-service/eval/prediction.py index 98f3eef1f..e5f3c8307 100644 --- a/wren-ai-service/eval/prediction.py +++ b/wren-ai-service/eval/prediction.py @@ -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], diff --git a/wren-ai-service/poetry.lock b/wren-ai-service/poetry.lock index 334f99294..bb15aaf48 100644 --- a/wren-ai-service/poetry.lock +++ b/wren-ai-service/poetry.lock @@ -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" diff --git a/wren-ai-service/pyproject.toml b/wren-ai-service/pyproject.toml index 4c8da1407..2121caa1c 100644 --- a/wren-ai-service/pyproject.toml +++ b/wren-ai-service/pyproject.toml @@ -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" diff --git a/wren-ai-service/src/__main__.py b/wren-ai-service/src/__main__.py index 26abf0049..8dec34a4c 100644 --- a/wren-ai-service/src/__main__.py +++ b/wren-ai-service/src/__main__.py @@ -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", diff --git a/wren-ai-service/src/core/pipeline.py b/wren-ai-service/src/core/pipeline.py index abefdbb95..f53096e00 100644 --- a/wren-ai-service/src/core/pipeline.py +++ b/wren-ai-service/src/core/pipeline.py @@ -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__) diff --git a/wren-ai-service/src/globals.py b/wren-ai-service/src/globals.py index 766b6d83e..ea337e746 100644 --- a/wren-ai-service/src/globals.py +++ b/wren-ai-service/src/globals.py @@ -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 diff --git a/wren-ai-service/src/pipelines/generation/followup_sql_generation.py b/wren-ai-service/src/pipelines/generation/followup_sql_generation.py index 9065230a0..773e4c0f2 100644 --- a/wren-ai-service/src/pipelines/generation/followup_sql_generation.py +++ b/wren-ai-service/src/pipelines/generation/followup_sql_generation.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_answer.py b/wren-ai-service/src/pipelines/generation/sql_answer.py index 15c61836e..c49f9d156 100644 --- a/wren-ai-service/src/pipelines/generation/sql_answer.py +++ b/wren-ai-service/src/pipelines/generation/sql_answer.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_breakdown.py b/wren-ai-service/src/pipelines/generation/sql_breakdown.py index 70a7e350c..a21c51037 100644 --- a/wren-ai-service/src/pipelines/generation/sql_breakdown.py +++ b/wren-ai-service/src/pipelines/generation/sql_breakdown.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_correction.py b/wren-ai-service/src/pipelines/generation/sql_correction.py index 8f2cd20e1..5462222fe 100644 --- a/wren-ai-service/src/pipelines/generation/sql_correction.py +++ b/wren-ai-service/src/pipelines/generation/sql_correction.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_expansion.py b/wren-ai-service/src/pipelines/generation/sql_expansion.py index b9dd1962e..2098d7b09 100644 --- a/wren-ai-service/src/pipelines/generation/sql_expansion.py +++ b/wren-ai-service/src/pipelines/generation/sql_expansion.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_explanation.py b/wren-ai-service/src/pipelines/generation/sql_explanation.py index 18d8c6c65..31139a38e 100644 --- a/wren-ai-service/src/pipelines/generation/sql_explanation.py +++ b/wren-ai-service/src/pipelines/generation/sql_explanation.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_generation.py b/wren-ai-service/src/pipelines/generation/sql_generation.py index 2fb70cfa5..e75f398e0 100644 --- a/wren-ai-service/src/pipelines/generation/sql_generation.py +++ b/wren-ai-service/src/pipelines/generation/sql_generation.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_regeneration.py b/wren-ai-service/src/pipelines/generation/sql_regeneration.py index 929469d73..1e63ef459 100644 --- a/wren-ai-service/src/pipelines/generation/sql_regeneration.py +++ b/wren-ai-service/src/pipelines/generation/sql_regeneration.py @@ -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() diff --git a/wren-ai-service/src/pipelines/generation/sql_summary.py b/wren-ai-service/src/pipelines/generation/sql_summary.py index 55dc49936..f6b3ed579 100644 --- a/wren-ai-service/src/pipelines/generation/sql_summary.py +++ b/wren-ai-service/src/pipelines/generation/sql_summary.py @@ -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() diff --git a/wren-ai-service/src/pipelines/indexing/indexing.py b/wren-ai-service/src/pipelines/indexing/indexing.py index 89741cb5d..878256760 100644 --- a/wren-ai-service/src/pipelines/indexing/indexing.py +++ b/wren-ai-service/src/pipelines/indexing/indexing.py @@ -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() diff --git a/wren-ai-service/src/pipelines/retrieval/historical_question.py b/wren-ai-service/src/pipelines/retrieval/historical_question.py index b9ad75241..39302e803 100644 --- a/wren-ai-service/src/pipelines/retrieval/historical_question.py +++ b/wren-ai-service/src/pipelines/retrieval/historical_question.py @@ -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") diff --git a/wren-ai-service/src/pipelines/retrieval/retrieval.py b/wren-ai-service/src/pipelines/retrieval/retrieval.py index e29dc5d41..647fb714d 100644 --- a/wren-ai-service/src/pipelines/retrieval/retrieval.py +++ b/wren-ai-service/src/pipelines/retrieval/retrieval.py @@ -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() diff --git a/wren-ai-service/src/providers/__init__.py b/wren-ai-service/src/providers/__init__.py index e69de29bb..2cdc9e9c7 100644 --- a/wren-ai-service/src/providers/__init__.py +++ b/wren-ai-service/src/providers/__init__.py @@ -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() + } diff --git a/wren-ai-service/src/providers/document_store/qdrant.py b/wren-ai-service/src/providers/document_store/qdrant.py index ca730533f..99054dba0 100644 --- a/wren-ai-service/src/providers/document_store/qdrant.py +++ b/wren-ai-service/src/providers/document_store/qdrant.py @@ -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 diff --git a/wren-ai-service/src/providers/embedder/azure_openai.py b/wren-ai-service/src/providers/embedder/azure_openai.py index 6f01edc56..8ed2dd452 100644 --- a/wren-ai-service/src/providers/embedder/azure_openai.py +++ b/wren-ai-service/src/providers/embedder/azure_openai.py @@ -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}") diff --git a/wren-ai-service/src/providers/embedder/ollama.py b/wren-ai-service/src/providers/embedder/ollama.py index e3a2fdd0b..47990a758 100644 --- a/wren-ai-service/src/providers/embedder/ollama.py +++ b/wren-ai-service/src/providers/embedder/ollama.py @@ -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) diff --git a/wren-ai-service/src/providers/embedder/openai.py b/wren-ai-service/src/providers/embedder/openai.py index d3ea5690b..44a926f98 100644 --- a/wren-ai-service/src/providers/embedder/openai.py +++ b/wren-ai-service/src/providers/embedder/openai.py @@ -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( diff --git a/wren-ai-service/src/providers/engine/wren.py b/wren-ai-service/src/providers/engine/wren.py index 0a6db9cb9..c8b217344 100644 --- a/wren-ai-service/src/providers/engine/wren.py +++ b/wren-ai-service/src/providers/engine/wren.py @@ -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") diff --git a/wren-ai-service/src/providers/llm/azure_openai.py b/wren-ai-service/src/providers/llm/azure_openai.py index 076a013b1..90a6016b2 100644 --- a/wren-ai-service/src/providers/llm/azure_openai.py +++ b/wren-ai-service/src/providers/llm/azure_openai.py @@ -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}") diff --git a/wren-ai-service/src/providers/llm/ollama.py b/wren-ai-service/src/providers/llm/ollama.py index 48de8b92d..dc852998c 100644 --- a/wren-ai-service/src/providers/llm/ollama.py +++ b/wren-ai-service/src/providers/llm/ollama.py @@ -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) diff --git a/wren-ai-service/src/providers/llm/openai.py b/wren-ai-service/src/providers/llm/openai.py index aa3b8c9bd..fb500b3b9 100644 --- a/wren-ai-service/src/providers/llm/openai.py +++ b/wren-ai-service/src/providers/llm/openai.py @@ -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}") diff --git a/wren-ai-service/src/utils.py b/wren-ai-service/src/utils.py index a8a095946..1274fb7b0 100644 --- a/wren-ai-service/src/utils.py +++ b/wren-ai-service/src/utils.py @@ -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"), diff --git a/wren-ai-service/tests/pytest/pipelines/test_ask.py b/wren-ai-service/tests/pytest/pipelines/test_ask.py index 33f1d8442..56435d48a 100644 --- a/wren-ai-service/tests/pytest/pipelines/test_ask.py +++ b/wren-ai-service/tests/pytest/pipelines/test_ask.py @@ -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 diff --git a/wren-ai-service/tests/pytest/pipelines/test_ask_details.py b/wren-ai-service/tests/pytest/pipelines/test_ask_details.py index b053cf538..784cf6dd9 100644 --- a/wren-ai-service/tests/pytest/pipelines/test_ask_details.py +++ b/wren-ai-service/tests/pytest/pipelines/test_ask_details.py @@ -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 diff --git a/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py b/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py index ea0b3ec69..59a26efcf 100644 --- a/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py +++ b/wren-ai-service/tests/pytest/pipelines/test_document_cleaner.py @@ -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 diff --git a/wren-ai-service/tests/pytest/services/test_ask.py b/wren-ai-service/tests/pytest/services/test_ask.py index 9d956c736..c05aaeaa0 100644 --- a/wren-ai-service/tests/pytest/services/test_ask.py +++ b/wren-ai-service/tests/pytest/services/test_ask.py @@ -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", } diff --git a/wren-ai-service/tests/pytest/services/test_ask_details.py b/wren-ai-service/tests/pytest/services/test_ask_details.py index 6920ae589..c4dd56aab 100644 --- a/wren-ai-service/tests/pytest/services/test_ask_details.py +++ b/wren-ai-service/tests/pytest/services/test_ask_details.py @@ -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", } diff --git a/wren-ai-service/tests/pytest/test_utils.py b/wren-ai-service/tests/pytest/test_utils.py index 2be8222ca..08ca6597b 100644 --- a/wren-ai-service/tests/pytest/test_utils.py +++ b/wren-ai-service/tests/pytest/test_utils.py @@ -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, }, )