diff --git a/.env.example b/.env.example index 67f1bd48..2ee155f0 100644 --- a/.env.example +++ b/.env.example @@ -23,11 +23,17 @@ ANTHROPIC_API_KEY=sk-YOUR-API-KEY-HERE OPENAI_API_KEY=sk-YOUR-API-KEY-HERE -PINECONE_KEY=YOUR-API-KEY-HERE +PINECONE_API_KEY=YOUR-API-KEY-HERE POSTGRES_URL=postgresql://localhost:5432/apollo_dev SENTRY_DSN=YOUR-API-KEY-HERE GITHUB_TOKEN=KEY +# Which backend serves docsite search reads: 'pinecone' (default) or 'postgres'. +DOCSITE_SEARCH_BACKEND=pinecone + +# Database for the Postgres docsite integration suite +POSTGRES_TEST_URL=postgresql://postgres:postgres@127.0.0.1:5433/postgres + # Langfuse observability LANGFUSE_SECRET_KEY=sk-lf-... LANGFUSE_PUBLIC_KEY=pk-lf-... diff --git a/poetry.lock b/poetry.lock index c413ba5d..3d85b154 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1724,6 +1724,114 @@ files = [ {file = "packaging-24.2.tar.gz", hash = "sha256:c228a6dc5e932d346bc5739379109d49e8853dd8223571c7c5b55260edc0b97f"}, ] +[[package]] +name = "pandas" +version = "2.3.3" +description = "Powerful data structures for data analysis, time series, and statistics" +optional = false +python-versions = ">=3.9" +groups = ["main"] +files = [ + {file = "pandas-2.3.3-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:376c6446ae31770764215a6c937f72d917f214b43560603cd60da6408f183b6c"}, + {file = "pandas-2.3.3-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:e19d192383eab2f4ceb30b412b22ea30690c9e618f78870357ae1d682912015a"}, + {file = "pandas-2.3.3-cp310-cp310-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5caf26f64126b6c7aec964f74266f435afef1c1b13da3b0636c7518a1fa3e2b1"}, + {file = "pandas-2.3.3-cp310-cp310-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:dd7478f1463441ae4ca7308a70e90b33470fa593429f9d4c578dd00d1fa78838"}, + {file = "pandas-2.3.3-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:4793891684806ae50d1288c9bae9330293ab4e083ccd1c5e383c34549c6e4250"}, + {file = "pandas-2.3.3-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:28083c648d9a99a5dd035ec125d42439c6c1c525098c58af0fc38dd1a7a1b3d4"}, + {file = "pandas-2.3.3-cp310-cp310-win_amd64.whl", hash = "sha256:503cf027cf9940d2ceaa1a93cfb5f8c8c7e6e90720a2850378f0b3f3b1e06826"}, + {file = "pandas-2.3.3-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:602b8615ebcc4a0c1751e71840428ddebeb142ec02c786e8ad6b1ce3c8dec523"}, + {file = "pandas-2.3.3-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:8fe25fc7b623b0ef6b5009149627e34d2a4657e880948ec3c840e9402e5c1b45"}, + {file = "pandas-2.3.3-cp311-cp311-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:b468d3dad6ff947df92dcb32ede5b7bd41a9b3cceef0a30ed925f6d01fb8fa66"}, + {file = "pandas-2.3.3-cp311-cp311-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b98560e98cb334799c0b07ca7967ac361a47326e9b4e5a7dfb5ab2b1c9d35a1b"}, + {file = "pandas-2.3.3-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:1d37b5848ba49824e5c30bedb9c830ab9b7751fd049bc7914533e01c65f79791"}, + {file = "pandas-2.3.3-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:db4301b2d1f926ae677a751eb2bd0e8c5f5319c9cb3f88b0becbbb0b07b34151"}, + {file = "pandas-2.3.3-cp311-cp311-win_amd64.whl", hash = "sha256:f086f6fe114e19d92014a1966f43a3e62285109afe874f067f5abbdcbb10e59c"}, + {file = "pandas-2.3.3-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:6d21f6d74eb1725c2efaa71a2bfc661a0689579b58e9c0ca58a739ff0b002b53"}, + {file = "pandas-2.3.3-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:3fd2f887589c7aa868e02632612ba39acb0b8948faf5cc58f0850e165bd46f35"}, + {file = "pandas-2.3.3-cp312-cp312-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ecaf1e12bdc03c86ad4a7ea848d66c685cb6851d807a26aa245ca3d2017a1908"}, + {file = "pandas-2.3.3-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b3d11d2fda7eb164ef27ffc14b4fcab16a80e1ce67e9f57e19ec0afaf715ba89"}, + {file = "pandas-2.3.3-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:a68e15f780eddf2b07d242e17a04aa187a7ee12b40b930bfdd78070556550e98"}, + {file = "pandas-2.3.3-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:371a4ab48e950033bcf52b6527eccb564f52dc826c02afd9a1bc0ab731bba084"}, + {file = "pandas-2.3.3-cp312-cp312-win_amd64.whl", hash = "sha256:a16dcec078a01eeef8ee61bf64074b4e524a2a3f4b3be9326420cabe59c4778b"}, + {file = "pandas-2.3.3-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:56851a737e3470de7fa88e6131f41281ed440d29a9268dcbf0002da5ac366713"}, + {file = "pandas-2.3.3-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:bdcd9d1167f4885211e401b3036c0c8d9e274eee67ea8d0758a256d60704cfe8"}, + {file = "pandas-2.3.3-cp313-cp313-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:e32e7cc9af0f1cc15548288a51a3b681cc2a219faa838e995f7dc53dbab1062d"}, + {file = "pandas-2.3.3-cp313-cp313-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:318d77e0e42a628c04dc56bcef4b40de67918f7041c2b061af1da41dcff670ac"}, + {file = "pandas-2.3.3-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:4e0a175408804d566144e170d0476b15d78458795bb18f1304fb94160cabf40c"}, + {file = "pandas-2.3.3-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:93c2d9ab0fc11822b5eece72ec9587e172f63cff87c00b062f6e37448ced4493"}, + {file = "pandas-2.3.3-cp313-cp313-win_amd64.whl", hash = "sha256:f8bfc0e12dc78f777f323f55c58649591b2cd0c43534e8355c51d3fede5f4dee"}, + {file = "pandas-2.3.3-cp313-cp313t-macosx_10_13_x86_64.whl", hash = "sha256:75ea25f9529fdec2d2e93a42c523962261e567d250b0013b16210e1d40d7c2e5"}, + {file = "pandas-2.3.3-cp313-cp313t-macosx_11_0_arm64.whl", hash = "sha256:74ecdf1d301e812db96a465a525952f4dde225fdb6d8e5a521d47e1f42041e21"}, + {file = "pandas-2.3.3-cp313-cp313t-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6435cb949cb34ec11cc9860246ccb2fdc9ecd742c12d3304989017d53f039a78"}, + {file = "pandas-2.3.3-cp313-cp313t-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:900f47d8f20860de523a1ac881c4c36d65efcb2eb850e6948140fa781736e110"}, + {file = "pandas-2.3.3-cp313-cp313t-musllinux_1_2_aarch64.whl", hash = "sha256:a45c765238e2ed7d7c608fc5bc4a6f88b642f2f01e70c0c23d2224dd21829d86"}, + {file = "pandas-2.3.3-cp313-cp313t-musllinux_1_2_x86_64.whl", hash = "sha256:c4fc4c21971a1a9f4bdb4c73978c7f7256caa3e62b323f70d6cb80db583350bc"}, + {file = "pandas-2.3.3-cp314-cp314-macosx_10_13_x86_64.whl", hash = "sha256:ee15f284898e7b246df8087fc82b87b01686f98ee67d85a17b7ab44143a3a9a0"}, + {file = "pandas-2.3.3-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:1611aedd912e1ff81ff41c745822980c49ce4a7907537be8692c8dbc31924593"}, + {file = "pandas-2.3.3-cp314-cp314-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6d2cefc361461662ac48810cb14365a365ce864afe85ef1f447ff5a1e99ea81c"}, + {file = "pandas-2.3.3-cp314-cp314-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:ee67acbbf05014ea6c763beb097e03cd629961c8a632075eeb34247120abcb4b"}, + {file = "pandas-2.3.3-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:c46467899aaa4da076d5abc11084634e2d197e9460643dd455ac3db5856b24d6"}, + {file = "pandas-2.3.3-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:6253c72c6a1d990a410bc7de641d34053364ef8bcd3126f7e7450125887dffe3"}, + {file = "pandas-2.3.3-cp314-cp314-win_amd64.whl", hash = "sha256:1b07204a219b3b7350abaae088f451860223a52cfb8a6c53358e7948735158e5"}, + {file = "pandas-2.3.3-cp314-cp314t-macosx_10_13_x86_64.whl", hash = "sha256:2462b1a365b6109d275250baaae7b760fd25c726aaca0054649286bcfbb3e8ec"}, + {file = "pandas-2.3.3-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:0242fe9a49aa8b4d78a4fa03acb397a58833ef6199e9aa40a95f027bb3a1b6e7"}, + {file = "pandas-2.3.3-cp314-cp314t-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a21d830e78df0a515db2b3d2f5570610f5e6bd2e27749770e8bb7b524b89b450"}, + {file = "pandas-2.3.3-cp314-cp314t-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:2e3ebdb170b5ef78f19bfb71b0dc5dc58775032361fa188e814959b74d726dd5"}, + {file = "pandas-2.3.3-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:d051c0e065b94b7a3cea50eb1ec32e912cd96dba41647eb24104b6c6c14c5788"}, + {file = "pandas-2.3.3-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:3869faf4bd07b3b66a9f462417d0ca3a9df29a9f6abd5d0d0dbab15dac7abe87"}, + {file = "pandas-2.3.3-cp39-cp39-macosx_10_9_x86_64.whl", hash = "sha256:c503ba5216814e295f40711470446bc3fd00f0faea8a086cbc688808e26f92a2"}, + {file = "pandas-2.3.3-cp39-cp39-macosx_11_0_arm64.whl", hash = "sha256:a637c5cdfa04b6d6e2ecedcb81fc52ffb0fd78ce2ebccc9ea964df9f658de8c8"}, + {file = "pandas-2.3.3-cp39-cp39-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:854d00d556406bffe66a4c0802f334c9ad5a96b4f1f868adf036a21b11ef13ff"}, + {file = "pandas-2.3.3-cp39-cp39-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:bf1f8a81d04ca90e32a0aceb819d34dbd378a98bf923b6398b9a3ec0bf44de29"}, + {file = "pandas-2.3.3-cp39-cp39-musllinux_1_2_aarch64.whl", hash = "sha256:23ebd657a4d38268c7dfbdf089fbc31ea709d82e4923c5ffd4fbd5747133ce73"}, + {file = "pandas-2.3.3-cp39-cp39-musllinux_1_2_x86_64.whl", hash = "sha256:5554c929ccc317d41a5e3d1234f3be588248e61f08a74dd17c9eabb535777dc9"}, + {file = "pandas-2.3.3-cp39-cp39-win_amd64.whl", hash = "sha256:d3e28b3e83862ccf4d85ff19cf8c20b2ae7e503881711ff2d534dc8f761131aa"}, + {file = "pandas-2.3.3.tar.gz", hash = "sha256:e05e1af93b977f7eafa636d043f9f94c7ee3ac81af99c13508215942e64c993b"}, +] + +[package.dependencies] +numpy = {version = ">=1.23.2", markers = "python_version == \"3.11\""} +python-dateutil = ">=2.8.2" +pytz = ">=2020.1" +tzdata = ">=2022.7" + +[package.extras] +all = ["PyQt5 (>=5.15.9)", "SQLAlchemy (>=2.0.0)", "adbc-driver-postgresql (>=0.8.0)", "adbc-driver-sqlite (>=0.8.0)", "beautifulsoup4 (>=4.11.2)", "bottleneck (>=1.3.6)", "dataframe-api-compat (>=0.1.7)", "fastparquet (>=2022.12.0)", "fsspec (>=2022.11.0)", "gcsfs (>=2022.11.0)", "html5lib (>=1.1)", "hypothesis (>=6.46.1)", "jinja2 (>=3.1.2)", "lxml (>=4.9.2)", "matplotlib (>=3.6.3)", "numba (>=0.56.4)", "numexpr (>=2.8.4)", "odfpy (>=1.4.1)", "openpyxl (>=3.1.0)", "pandas-gbq (>=0.19.0)", "psycopg2 (>=2.9.6)", "pyarrow (>=10.0.1)", "pymysql (>=1.0.2)", "pyreadstat (>=1.2.0)", "pytest (>=7.3.2)", "pytest-xdist (>=2.2.0)", "python-calamine (>=0.1.7)", "pyxlsb (>=1.0.10)", "qtpy (>=2.3.0)", "s3fs (>=2022.11.0)", "scipy (>=1.10.0)", "tables (>=3.8.0)", "tabulate (>=0.9.0)", "xarray (>=2022.12.0)", "xlrd (>=2.0.1)", "xlsxwriter (>=3.0.5)", "zstandard (>=0.19.0)"] +aws = ["s3fs (>=2022.11.0)"] +clipboard = ["PyQt5 (>=5.15.9)", "qtpy (>=2.3.0)"] +compression = ["zstandard (>=0.19.0)"] +computation = ["scipy (>=1.10.0)", "xarray (>=2022.12.0)"] +consortium-standard = ["dataframe-api-compat (>=0.1.7)"] +excel = ["odfpy (>=1.4.1)", "openpyxl (>=3.1.0)", "python-calamine (>=0.1.7)", "pyxlsb (>=1.0.10)", "xlrd (>=2.0.1)", "xlsxwriter (>=3.0.5)"] +feather = ["pyarrow (>=10.0.1)"] +fss = ["fsspec (>=2022.11.0)"] +gcp = ["gcsfs (>=2022.11.0)", "pandas-gbq (>=0.19.0)"] +hdf5 = ["tables (>=3.8.0)"] +html = ["beautifulsoup4 (>=4.11.2)", "html5lib (>=1.1)", "lxml (>=4.9.2)"] +mysql = ["SQLAlchemy (>=2.0.0)", "pymysql (>=1.0.2)"] +output-formatting = ["jinja2 (>=3.1.2)", "tabulate (>=0.9.0)"] +parquet = ["pyarrow (>=10.0.1)"] +performance = ["bottleneck (>=1.3.6)", "numba (>=0.56.4)", "numexpr (>=2.8.4)"] +plot = ["matplotlib (>=3.6.3)"] +postgresql = ["SQLAlchemy (>=2.0.0)", "adbc-driver-postgresql (>=0.8.0)", "psycopg2 (>=2.9.6)"] +pyarrow = ["pyarrow (>=10.0.1)"] +spss = ["pyreadstat (>=1.2.0)"] +sql-other = ["SQLAlchemy (>=2.0.0)", "adbc-driver-postgresql (>=0.8.0)", "adbc-driver-sqlite (>=0.8.0)"] +test = ["hypothesis (>=6.46.1)", "pytest (>=7.3.2)", "pytest-xdist (>=2.2.0)"] +xml = ["lxml (>=4.9.2)"] + +[[package]] +name = "pgvector" +version = "0.5.0" +description = "pgvector support for Python" +optional = false +python-versions = ">=3.10" +groups = ["main"] +files = [ + {file = "pgvector-0.5.0-py3-none-any.whl", hash = "sha256:fedc9800894e6da2be51358d7b7c574bf34f247ca741a5a09513622135f5964f"}, + {file = "pgvector-0.5.0.tar.gz", hash = "sha256:07a9dcf735696879406983afc6eba9a787cef7c0cf6c367ca1a5779f036dee74"}, +] + [[package]] name = "pinecone" version = "7.3.0" @@ -2269,6 +2377,18 @@ files = [ [package.extras] cli = ["click (>=5.0)"] +[[package]] +name = "pytz" +version = "2026.3.post1" +description = "World timezone definitions, modern and historical" +optional = false +python-versions = "*" +groups = ["main"] +files = [ + {file = "pytz-2026.3.post1-py2.py3-none-any.whl", hash = "sha256:dd95840dd199baea12d9cc096a1d452caa6596a1c1e4b5f3dbd1541855d5e815"}, + {file = "pytz-2026.3.post1.tar.gz", hash = "sha256:2211d3fcf9a797d3405cac96ac7f61d80e6a644f72a3309607282fe8a2010c5d"}, +] + [[package]] name = "pyyaml" version = "6.0.3" @@ -2958,6 +3078,18 @@ files = [ [package.dependencies] typing-extensions = ">=4.12.0" +[[package]] +name = "tzdata" +version = "2026.3" +description = "Provider of IANA time zone data" +optional = false +python-versions = ">=2" +groups = ["main"] +files = [ + {file = "tzdata-2026.3-py2.py3-none-any.whl", hash = "sha256:dc096730c87af6cab1b171c9d532be840741ff5d459015e7f6947bd7d7e54931"}, + {file = "tzdata-2026.3.tar.gz", hash = "sha256:4a1518b8993086a7982523e071643f3c0e5f213e75b21318e78bcabfff9d1415"}, +] + [[package]] name = "urllib3" version = "2.7.0" @@ -3625,4 +3757,4 @@ cffi = ["cffi (>=1.17,<2.0) ; platform_python_implementation != \"PyPy\" and pyt [metadata] lock-version = "2.1" python-versions = "3.11.*" -content-hash = "6551e8fae076b04ea022e47a09446e6573992a90751c4ce633f9e2e0d73f035e" +content-hash = "e0f7dc02320564ac4dc7929bf5ff747716b4ab90810f4acc421db052f49014cc" diff --git a/pyproject.toml b/pyproject.toml index f67ccac0..04470372 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,6 +26,8 @@ psycopg2-binary = "^2.9.10" langfuse = "^4.14.1" opentelemetry-instrumentation-anthropic = "^0.62.1" opentelemetry-instrumentation-threading = "0.65b0" +pgvector = "^0.5.0" +pandas = "^2.2" [tool.poetry.group.dev] optional = false @@ -50,6 +52,7 @@ testpaths = [ "services/job_chat/tests", "services/latest_adaptors/tests", "services/search_docsite/tests", + "services/embed_docsite/tests", "services/tools", ] diff --git a/services/db_migrations.py b/services/db_migrations.py new file mode 100644 index 00000000..47e09f74 --- /dev/null +++ b/services/db_migrations.py @@ -0,0 +1,64 @@ +"""Versioned schema migrations for the Python-owned docs database (POSTGRES_URL). + +Applies .sql files in lexical order, records applied filenames so re-runs are a +no-op, and takes an advisory lock so concurrent starters queue. + +The tracking table (_migrations_docs) and lock key are distinct from the +TypeScript runner's in platform/src/db/migrate.ts, because APOLLO_CLIENTS_DB_URL +falls back to POSTGRES_URL locally and both runners can target one database. +""" + +from pathlib import Path + +from util import create_logger + +logger = create_logger("db_migrations") + +MIGRATIONS_DIR = Path(__file__).parent / "migrations" + +# Distinct from the TypeScript runner's 8314_2025 so the two never block each other. +MIGRATION_LOCK_KEY = 8314_2026 + +CREATE_TRACKING_TABLE_SQL = """ +CREATE TABLE IF NOT EXISTS _migrations_docs ( + filename TEXT PRIMARY KEY, + applied_at TIMESTAMPTZ NOT NULL DEFAULT now() +) +""" + + +def _migration_files(): + """Every .sql file in the migrations directory, in lexical order.""" + if not MIGRATIONS_DIR.is_dir(): + return [] + return sorted(MIGRATIONS_DIR.glob("*.sql")) + + +def run_migrations(conn): + """Apply any migrations not yet recorded. Returns the count applied this run. + + Everything happens in one transaction: the advisory lock is held for its + duration, so a racing process waits and then sees the migrations already + recorded rather than colliding on CREATE TABLE. + """ + files = _migration_files() + + with conn.cursor() as cur: + cur.execute("SELECT pg_advisory_xact_lock(%s)", (MIGRATION_LOCK_KEY,)) + cur.execute(CREATE_TRACKING_TABLE_SQL) + + cur.execute("SELECT filename FROM _migrations_docs") + already_applied = {row[0] for row in cur.fetchall()} + + pending = [f for f in files if f.name not in already_applied] + for path in pending: + logger.info(f"Applying migration {path.name}") + cur.execute(path.read_text(encoding="utf-8")) + cur.execute("INSERT INTO _migrations_docs (filename) VALUES (%s)", (path.name,)) + + conn.commit() + + if pending: + logger.info(f"Applied {len(pending)} migration(s)") + + return len(pending) diff --git a/services/embed_docsite/README.md b/services/embed_docsite/README.md index 6129f303..71512072 100644 --- a/services/embed_docsite/README.md +++ b/services/embed_docsite/README.md @@ -1,10 +1,13 @@ ## Embed Docsite (RAG) -This service embeds the OpenFn Documentation to a vector database. It downloads, chunks, processes metadata, embeds and uploads the documentation to a vector database (Pinecone). +This service embeds the OpenFn Documentation to a vector database. It downloads, +chunks, processes metadata, embeds and uploads the documentation to a vector +database (Pinecone). ## Usage - Embedding OpenFn Documentation -The vector database used here is Pinecone. To obtain the env variables follow these steps: +The vector database used here is Pinecone. To obtain the env variables follow +these steps: 1. Create an account on [Pinecone] and set up a free cluster. 2. Obtain the URL and token for the cluster and add them to the `.env` file. @@ -15,6 +18,7 @@ The vector database used here is Pinecone. To obtain the env variables follow th ```bash openfn apollo embed_docsite tmp/payload.json ``` + To run directly from this repo (note that the server must be started): ```bash @@ -22,19 +26,35 @@ bun py embed_docsite tmp/payload.json -O ``` ## Implementation -The service uses the DocsiteProcessor to download the documentation and chunk it into smaller parts. The DocsiteIndexer formats metadata, creates a new collection, embeds the chunked texts (OpenAI) and uploads them into the vector database (Pinecone). + +The service uses the DocsiteProcessor to download the documentation and chunk it +into smaller parts. The DocsiteIndexer formats metadata, creates a new +collection, embeds the chunked texts (OpenAI) and uploads them into the vector +database (Pinecone). The chunked texts can be viewed in `tmp/split_sections`. ## Payload Reference + +The write target is independent of the read backend (`DOCSITE_SEARCH_BACKEND`), +so a Postgres batch can be built while Pinecone still serves search traffic. + The input payload is a JSON object. All parameters are optional: ```js { + "target": "pinecone", // 'pinecone' | 'postgres'. Defaults to pinecone. Chooses the write destination. "docs_to_upload": ["adaptor_docs", "general_docs", "adaptor_functions"], // Select from 3 types of documentation to upload - "collection_name": "docsite-20250225", // Name of the collection in the vector database (defaults to the current date) - "index_name": "docsite", // Name of the index in the vector database (an index contains collections; defaults to docsite) "docs_to_ignore": ["job-examples.md", "release-notes.md"], // Titles of documents that should not be indexed - "max_total_collections" : 3 // The max number of collections to keep in the vector database. This will delete older collections by date. + "chunk_target_length": 1000, // Target chunk size in characters + "chunk_min_length": 700, // Minimum chunk size before merging with the next split + + // Pinecone target only: + "collection_name": "docsite-20250225", // Namespace (defaults to the current timestamp) + "index_name": "docsite", // Name of the index in the vector database (an index contains collections; defaults to docsite) + "max_total_collections": 3, // The max number of collections to keep in the vector database. This will delete older collections by date. + + // Postgres target only: + "keep_batches": 2 // Number of recent complete batches to retain when pruning } ``` diff --git a/services/embed_docsite/__init__.py b/services/embed_docsite/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/embed_docsite/docsite_indexer.py b/services/embed_docsite/docsite_indexer.py index 0e418948..c7a32480 100644 --- a/services/embed_docsite/docsite_indexer.py +++ b/services/embed_docsite/docsite_indexer.py @@ -1,178 +1,199 @@ -import os -import time -from datetime import datetime -import pandas as pd -from pinecone import Pinecone, ServerlessSpec -from langchain_pinecone import PineconeVectorStore +from contextlib import contextmanager + from langchain_openai import OpenAIEmbeddings -from langchain_community.document_loaders import DataFrameLoader -from util import create_logger, ApolloError +from pgvector import Vector +from pgvector.psycopg2 import register_vector +from psycopg2.extras import execute_values +from util import create_logger logger = create_logger("DocsiteIndexer") -class DocsiteIndexer: - """ - Initialize vectorstore and insert new documents. Create a new index if needed. +ALL_DOCS_TYPES = ["adaptor_docs", "general_docs", "adaptor_functions"] - :param collection_name: Vectorstore collection name (namespace) to store documents - :param index_name: Vectorstore index name (default: docsite) - :param embeddings: LangChain embedding type (default: OpenAIEmbeddings()) - :param dimension: Embedding dimension (default: 1536 for OpenAI Embeddings) - :param max_total_collections: Max total collections in index. Delete old collections by date if exceeded after a new upload (default: 50) + +@contextmanager +def autocommit(conn): + """Run DDL that cannot execute inside a transaction (CREATE/DROP INDEX + CONCURRENTLY). + + Commits first: psycopg2 refuses to change autocommit while a transaction is + open, and a preceding SELECT is enough to open one. """ - def __init__(self, collection_name=None, index_name="docsite", embeddings=OpenAIEmbeddings(), dimension=1536, max_total_collections=50): - self.collection_name = collection_name if collection_name is not None else f"docsite-{datetime.now().strftime('%Y%m%d%H%M')}" - self.index_name = index_name - self.embeddings = embeddings - self.dimension = dimension - self.max_total_collections = max_total_collections - self.pc = Pinecone(api_key=os.environ.get("PINECONE_API_KEY")) + previous = conn.autocommit + conn.commit() + conn.autocommit = True + try: + yield + finally: + conn.autocommit = previous - if not self.index_exists(): - self.create_index() - self.index = self.pc.Index(self.index_name) - self.vectorstore = PineconeVectorStore(index_name=index_name, namespace=self.collection_name, embedding=embeddings) +def register_vector_type(conn): + """Register pgvector's psycopg2 adapters on this connection.""" + register_vector(conn) - def insert_documents(self, inputs, metadata_dict): - """ - Create the index if it does not exist and insert the input documents. - - :param inputs: Dictionary containing name, docs_type, and doc_chunk - :param metadata_dict: Metadata dict with document titles as keys (from DocsiteProcessor) - :return: Initialized indices - """ - # Get vector count before insertion for verification - try: - stats = self.index.describe_index_stats() - vectors_before = stats.namespaces.get(self.collection_name, {}).get("vector_count", 0) - logger.info(f"Current vector count in namespace '{self.collection_name}': {vectors_before}") - except Exception as e: - logger.warning(f"Could not get vector count before insertion: {str(e)}") - vectors_before = 0 - - df = self.preprocess_metadata(inputs=inputs, metadata_dict=metadata_dict) - logger.info(f"Input metadata preprocessed") - loader = DataFrameLoader(df, page_content_column="text") - docs = loader.load() - logger.info(f"Inputs processed into LangChain docs") - logger.info(f"Uploading {len(docs)} documents to index...") - - idx = self.vectorstore.add_documents( - documents=docs - ) - sleep_time = 10 - max_wait_time = 150 - elapsed_time = 0 - logger.info(f"Waiting up to {max_wait_time}s to verify upload count") - - while elapsed_time < max_wait_time: - time.sleep(sleep_time) - elapsed_time += sleep_time - - # Verify the upload by checking the vector count - try: - stats = self.index.describe_index_stats() - vectors_after = stats.namespaces.get(self.collection_name, {}).get("vector_count", 0) - logger.info(f"Vector count after {elapsed_time}s: {vectors_after}") - - if vectors_after >= vectors_before + len(docs): - logger.info(f"Successfully added {vectors_after - vectors_before} vectors to namespace '{self.collection_name}'") - break - else: - logger.warning(f"No new vectors were added to namespace '{self.collection_name}' after {sleep_time}s") - except Exception as e: - logger.warning(f"Could not verify vector insertion: {str(e)}") - - if vectors_after <= vectors_before: - logger.warning(f"Could not verify full dataset upload to namespace '{self.collection_name}' after {max_wait_time}s") - - self.delete_old_collections(self.max_total_collections) - - return idx - - def delete_collection(self): +class DocsiteIndexer: + """ + Builds versioned "batches" of embedded docsite chunks in Postgres. + + A batch is a full, self-consistent snapshot across all docs_types. Batches + are built invisibly (status='building'), then promoted to 'complete' + atomically, so readers only ever see a finished batch. + + :param chunk_target_length: Target chunk size in characters (default: 1000) + :param chunk_min_length: Minimum chunk size before merging with the next split (default: 700) + :param keep_batches: Number of most-recent complete batches to retain when pruning (default: 2) + """ + + def __init__(self, chunk_target_length=1000, chunk_min_length=700, keep_batches=2): + self.chunk_target_length = chunk_target_length + self.chunk_min_length = chunk_min_length + self.keep_batches = keep_batches + self._embeddings = None + + @property + def embeddings(self): + """Lazily construct the OpenAI embeddings client.""" + if self._embeddings is None: + self._embeddings = OpenAIEmbeddings() + return self._embeddings + + def start_batch(self, conn, docs_types): + """Insert a new 'building' batch row and return its id.""" + sql = """ + INSERT INTO docsite_batches (status, docs_types, chunk_target_length, chunk_min_length, embedding_model) + VALUES ('building', %s, %s, %s, %s) + RETURNING id """ - Deletes the entire collection (namespace) and all its contents. - This operation cannot be undone and removes both the collection structure and all vectors/documents within it. + with conn.cursor() as cur: + cur.execute(sql, (docs_types, self.chunk_target_length, self.chunk_min_length, self.embeddings.model)) + batch_id = cur.fetchone()[0] + conn.commit() + logger.info(f"Started batch {batch_id} for docs_types={docs_types}") + return batch_id + + def insert_documents(self, conn, batch_id, documents, metadata_dict): + """Embed and bulk-insert chunks for this batch. Returns the number of chunks inserted.""" + if not documents: + return 0 + + texts = [doc["doc_chunk"] for doc in documents] + embeddings = self._embed_in_batches(texts) + + doc_title_indices = {} + rows = [] + for doc, embedding in zip(documents, embeddings, strict=True): + doc_title = doc["name"].removesuffix(".md") + chunk_index = doc_title_indices.get(doc_title, 0) + doc_title_indices[doc_title] = chunk_index + 1 + rows.append((batch_id, doc_title, doc["docs_type"], chunk_index, doc["doc_chunk"], Vector(embedding))) + + insert_sql = """ + INSERT INTO docsite_chunks (batch_id, doc_title, docs_type, chunk_index, text, embedding) + VALUES %s """ - self.index.delete(delete_all=True, namespace=self.collection_name) - - def delete_old_collections(self, max_total_collections): - """Retrieve docsite uploads by collection name from Pinecone and delete them if there are more than max_total_collections.""" - - logger.info(f"Fetching outdated docsite collections") - pc = Pinecone(api_key=os.environ.get("PINECONE_API_KEY")) - index = pc.Index("docsite") - index_stats = index.describe_index_stats() - namespaces = index_stats.get('namespaces', {}).keys() - valid_namespaces = sorted( - (ns for ns in namespaces if ns.startswith("docsite-") and ns[8:].isdigit() and len(ns) == 16), - reverse=False + with conn.cursor() as cur: + execute_values(cur, insert_sql, rows) + conn.commit() + + logger.info(f"Inserted {len(rows)} chunks into batch {batch_id}") + return len(rows) + + def _embed_in_batches(self, texts, batch_size=100): + """Call the OpenAI embeddings API in batches of batch_size texts.""" + embeddings = [] + for i in range(0, len(texts), batch_size): + batch = texts[i:i + batch_size] + embeddings.extend(self.embeddings.embed_documents(batch)) + return embeddings + + def copy_forward_missing_docs_types(self, conn, batch_id, docs_types_present): + """Copy chunks for docs_types not in this run from the previous complete batch, + so every complete batch is a full snapshot across all docs_types. Returns rows copied.""" + missing_types = [t for t in ALL_DOCS_TYPES if t not in docs_types_present] + if not missing_types: + return 0 + + with conn.cursor() as cur: + cur.execute("SELECT id FROM docsite_batches WHERE status = 'complete' ORDER BY id DESC LIMIT 1") + row = cur.fetchone() + if row is None: + logger.info("No previous complete batch to copy forward from") + return 0 + previous_batch_id = row[0] + + cur.execute( + """ + INSERT INTO docsite_chunks (batch_id, doc_title, docs_type, chunk_index, text, embedding) + SELECT %s, doc_title, docs_type, chunk_index, text, embedding + FROM docsite_chunks + WHERE batch_id = %s AND docs_type = ANY(%s) + """, + (batch_id, previous_batch_id, missing_types), ) - if len(valid_namespaces) > max_total_collections: - logger.info(f"Deleting outdated docsite collections") - for old_collection in valid_namespaces[:max_total_collections]: - self.index.delete(delete_all=True, namespace=old_collection) - logger.info(f"Deleted collection {old_collection}") - - if not valid_namespaces: - logger.info(f"No valid namespaces found in the index when deleting old collections.") - - def create_index(self): - """Creates a new Pinecone index if it does not exist.""" - - if not self.index_exists(): - self.pc.create_index( - name=self.index_name, - dimension=self.dimension, - metric="cosine", - spec=ServerlessSpec(cloud="aws", region="us-east-1") + copied = cur.rowcount + conn.commit() + + logger.info(f"Copied {copied} chunks forward for docs_types={missing_types} from batch {previous_batch_id}") + return copied + + def build_index(self, conn, batch_id): + """Build a per-batch partial HNSW index.""" + with autocommit(conn): + with conn.cursor() as cur: + cur.execute( + f""" + CREATE INDEX CONCURRENTLY IF NOT EXISTS idx_docsite_chunks_hnsw_{batch_id} + ON docsite_chunks USING hnsw (embedding vector_cosine_ops) + WITH (m = 16, ef_construction = 64) + WHERE batch_id = {batch_id} + """ + ) + logger.info(f"Built HNSW index for batch {batch_id}") + + def promote_batch(self, conn, batch_id, chunk_count): + """Flip a batch to 'complete' the moment it becomes visible to readers.""" + with conn.cursor() as cur: + cur.execute( + "UPDATE docsite_batches SET status = 'complete', completed_at = now(), chunk_count = %s WHERE id = %s", + (chunk_count, batch_id), ) - while not self.pc.describe_index(self.index_name).status["ready"]: - time.sleep(1) + conn.commit() + logger.info(f"Promoted batch {batch_id} ({chunk_count} chunks)") - def index_exists(self): - """Check if the index exists in Pinecone.""" - existing_indexes = [index_info["name"] for index_info in self.pc.list_indexes()] + def fail_batch(self, conn, batch_id): + """Mark a batch 'failed' after an aborted build. - return self.index_name in existing_indexes - - def preprocess_metadata(self, inputs, page_content_column="text", add_chunk_as_metadata=False, metadata_cols=None, metadata_dict=None): - """ - Create a DataFrame for indexing from input documents and metadata. - - :param inputs: Dictionary containing name, docs_type, and doc_chunk - :param page_content_column: Name of the field which will be embedded (default: text) - :param add_chunk_as_metadata: Copy the text to embed as a separate metadata field (default: False) - :param metadata_cols: Optional list of metadata columns to include (default: None) - :param metadata_dict: Dictionary mapping names to metadata dictionaries (default: None) - :return: pandas.DataFrame with text and metadata columns + Rolls back first: the connection is in an aborted transaction from + whatever error got us here, so any statement would raise + InFailedSqlTransaction. """ - - # Create DataFrame from the inputs (doc_chunk, name, docs_type) - df = pd.DataFrame(inputs) - - # Rename some columns for metadata upload - df = df.rename(columns={"doc_chunk": page_content_column, "name": "doc_title"}) - - df["doc_title"] = df["doc_title"].str.replace(".md$", "", regex=True) - - # Optionally add chunk to metadata for keyword searching - if add_chunk_as_metadata: - df["embedding_text"] = df[page_content_column] - - # Add further metadata columns if specified - if metadata_cols: - for col in metadata_cols: - df[col] = metadata_dict.get(inputs["name"], {}).get(col) - - return df - - - - - - - \ No newline at end of file + conn.rollback() + with conn.cursor() as cur: + cur.execute("UPDATE docsite_batches SET status = 'failed' WHERE id = %s", (batch_id,)) + conn.commit() + logger.info(f"Marked batch {batch_id} failed") + + def prune_old_batches(self, conn, keep_batches=None): + """Delete complete batches older than the newest `keep_batches`, dropping their + partial indexes first. Returns the list of pruned batch ids.""" + keep = keep_batches if keep_batches is not None else self.keep_batches + + with conn.cursor() as cur: + cur.execute( + "SELECT id FROM docsite_batches WHERE status = 'complete' ORDER BY id DESC OFFSET %s", + (keep,), + ) + old_batch_ids = [row[0] for row in cur.fetchall()] + + for batch_id in old_batch_ids: + with autocommit(conn): + with conn.cursor() as cur: + cur.execute(f"DROP INDEX CONCURRENTLY IF EXISTS idx_docsite_chunks_hnsw_{batch_id}") + with conn.cursor() as cur: + cur.execute("DELETE FROM docsite_batches WHERE id = %s", (batch_id,)) + conn.commit() + logger.info(f"Pruned batch {batch_id}") + + return old_batch_ids diff --git a/services/embed_docsite/docsite_processor.py b/services/embed_docsite/docsite_processor.py index 63ec3d85..5a14e395 100644 --- a/services/embed_docsite/docsite_processor.py +++ b/services/embed_docsite/docsite_processor.py @@ -1,13 +1,15 @@ import json import os -import logging import re -import requests + import nltk from embed_docsite.github_utils import get_docs -from util import create_logger, ApolloError +from util import create_logger -nltk.download('punkt_tab') +try: + nltk.data.find('tokenizers/punkt_tab') +except LookupError: + nltk.download('punkt_tab', quiet=True) logger = create_logger("DocsiteProcessor") @@ -18,10 +20,13 @@ class DocsiteProcessor: :param docs_type: Type of documentation being processed ("adaptor_functions", "general_docs", "adaptor_docs") :param output_dir: Directory to store processed chunks (default: "./tmp/split_sections"). """ - def __init__(self, docs_type, docs_to_ignore=["job-examples.md", "release-notes.md"], output_dir="./tmp/split_sections"): + def __init__(self, docs_type, docs_to_ignore=["job-examples.md", "release-notes.md"], target_length=1000, min_length=700, overlap=1, output_dir="./tmp/split_sections"): self.output_dir = output_dir self.docs_type = docs_type self.docs_to_ignore = docs_to_ignore + self.target_length = target_length + self.min_length = min_length + self.overlap = overlap self.metadata_dict = None def get_preprocessed_docs(self): @@ -41,7 +46,7 @@ def get_preprocessed_docs(self): return chunks, metadata_dict - def _chunk_adaptor_docs(self, json_data, target_length=1000, overlap=1, min_length=700): + def _chunk_adaptor_docs(self, json_data): """Extract and clean docs from adaptor data, and chunk according to a target and minimum chunk sizes.""" output = [] metadata_dict = dict() @@ -68,13 +73,12 @@ def _chunk_adaptor_docs(self, json_data, target_length=1000, overlap=1, min_leng # Split by headers, and where needed, sentences splits = self._split_by_headers(docs) - splits = self._split_oversized_chunks(chunks=splits, target_length=target_length) - chunks = self._accumulate_chunks(splits=splits, target_length=target_length, overlap=overlap, min_length=min_length) + splits = self._split_oversized_chunks(chunks=splits, target_length=self.target_length) + chunks = self._accumulate_chunks(splits=splits, target_length=self.target_length, overlap=self.overlap, min_length=self.min_length) for chunk in chunks: output.append({"name": name, "docs_type": self.docs_type, "doc_chunk": chunk}) - # self.metadata_dict = metadata_dict self._write_chunks_to_file(chunks=output, file_name=f"{self.docs_type}_chunks.json") return output, metadata_dict @@ -136,7 +140,7 @@ def _accumulate_chunks(self, splits, target_length, overlap, min_length): if len(current_chunk) >= min_length: accumulated.append(current_chunk) # Store the completed chunk - # add overlap + # Add overlap if self.docs_type == "adaptor_functions": overlap_sections = " ".join(current_chunk.split("\n")[-overlap:]) else: @@ -176,4 +180,4 @@ def _write_chunks_to_file(self, chunks, file_name): with open(output_file, 'w') as f: json.dump(chunks, f, indent=2) - logger.info(f"Content written to {output_file}") \ No newline at end of file + logger.info(f"Content written to {output_file}") diff --git a/services/embed_docsite/embed_docsite.py b/services/embed_docsite/embed_docsite.py index 72c6887d..6e03c9c3 100644 --- a/services/embed_docsite/embed_docsite.py +++ b/services/embed_docsite/embed_docsite.py @@ -1,65 +1,157 @@ import os -import json + +from db_migrations import run_migrations from dotenv import load_dotenv -import pandas as pd -from util import create_logger, ApolloError +from embed_docsite.docsite_indexer import ( + ALL_DOCS_TYPES, + DocsiteIndexer, + register_vector_type, +) from embed_docsite.docsite_processor import DocsiteProcessor -from embed_docsite.docsite_indexer import DocsiteIndexer +from embed_docsite.pinecone_legacy_indexer import LegacyPineconeDocsiteIndexer +from util import ApolloError, create_logger, get_db_connection logger = create_logger("embed_docsite") -def main(data): +VALID_TARGETS = ("pinecone", "postgres") + + +def _collect_documents(docs_to_upload, docs_to_ignore, chunk_target_length, chunk_min_length): + """Download and chunk every requested docs_type. Shared by both targets.""" + documents = [] + metadata_dict = {} + for docs_type in docs_to_upload: + processor = DocsiteProcessor( + docs_type=docs_type, + docs_to_ignore=docs_to_ignore, + target_length=chunk_target_length, + min_length=chunk_min_length, + ) + type_documents, type_metadata = processor.get_preprocessed_docs() + documents.extend(type_documents) + metadata_dict.update(type_metadata) + return documents, metadata_dict + + +def main(data: dict) -> dict: logger.info("Starting...") - # Get selection of doc types to upload, or default to all - docs_to_upload = data.get("docs_to_upload", ["adaptor_docs", "general_docs", "adaptor_functions"]) + target = data.get("target", "pinecone") + docs_to_upload = data.get("docs_to_upload", ALL_DOCS_TYPES) docs_to_ignore = data.get("docs_to_ignore", ["job-examples.md", "release-notes.md"]) + chunk_target_length = data.get("chunk_target_length", 1000) + chunk_min_length = data.get("chunk_min_length", 700) + keep_batches = data.get("keep_batches", 2) - # Get other fields - index_params = {} - index_param_options = ["collection_name", "index_name", "max_total_collections"] + if target not in VALID_TARGETS: + raise ApolloError( + 400, + f"Unknown target '{target}'. Expected 'pinecone' or 'postgres'", + type="BAD_REQUEST", + ) - for key in index_param_options: - if key in data: - index_params[key] = data[key] - - # Set API keys load_dotenv(override=True) - if data.get("PINECONE_API_KEY", ""): - PINECONE_API_KEY = data["PINECONE_API_KEY"] - else: - PINECONE_API_KEY = os.environ.get("PINECONE_API_KEY") - - if data.get("OPENAI_API_KEY", ""): - OPENAI_API_KEY = data["OPENAI_API_KEY"] - else: - OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY") - - # Check for missing keys - missing_keys = [] - - if not OPENAI_API_KEY: - missing_keys.append("OPENAI_API_KEY") - if not PINECONE_API_KEY: - missing_keys.append("PINECONE_API_KEY") - - if missing_keys: - msg = f'Missing API keys: {", ".join(missing_keys)}' + openai_api_key = data.get("OPENAI_API_KEY") or os.environ.get("OPENAI_API_KEY") + if not openai_api_key: + msg = "Missing API key: OPENAI_API_KEY" logger.error(msg) - raise ApolloError(500, f'Missing API keys: {", ".join(missing_keys)}. Add to payload or environment.', type="BAD_REQUEST") + raise ApolloError(500, f"{msg}. Add to payload or environment", type="BAD_REQUEST") - # Initialize indexer - docsite_indexer = DocsiteIndexer(**(index_params or {})) + documents, metadata_dict = _collect_documents( + docs_to_upload, docs_to_ignore, chunk_target_length, chunk_min_length, + ) - # Add docs - for docs_type in docs_to_upload: - # Download and process - docsite_processor = DocsiteProcessor(docs_type=docs_type, docs_to_ignore=docs_to_ignore) - documents, metadata_dict = docsite_processor.get_preprocessed_docs() + if target == "pinecone": + return _upload_to_pinecone(data, documents, metadata_dict, docs_to_upload) + + return _upload_to_postgres( + documents, metadata_dict, docs_to_upload, chunk_target_length, chunk_min_length, keep_batches, + ) + + +def _upload_to_pinecone(data, documents, metadata_dict, docs_to_upload): + """Legacy write path to pinecone.""" + pinecone_api_key = data.get("PINECONE_API_KEY") or os.environ.get("PINECONE_API_KEY") + if not pinecone_api_key: + msg = "Missing API key: PINECONE_API_KEY" + logger.error(msg) + raise ApolloError(500, f"{msg}. Add to payload or environment", type="BAD_REQUEST") + + index_params = { + key: data[key] + for key in ("collection_name", "index_name", "max_total_collections") + if key in data + } + indexer = LegacyPineconeDocsiteIndexer(**index_params) + indexer.insert_documents(documents, metadata_dict) + + return { + "target": "pinecone", + "collection_name": indexer.collection_name, + "docs_types": docs_to_upload, + "chunk_count": len(documents), + } + + +def _mark_failed(indexer, conn, batch_id): + """Best-effort bookkeeping: never let it replace the error that caused it.""" + try: + indexer.fail_batch(conn, batch_id) + except Exception as exc: + logger.error(f"Could not mark batch {batch_id} failed: {exc}") + + +def _prune_old_batches(indexer, conn): + """Best-effort cleanup of older, unrelated batches. Runs after the new + batch is already promoted.""" + try: + return indexer.prune_old_batches(conn) + except Exception as exc: + logger.error(f"Could not prune old batches: {exc}") + return [] + + +def _upload_to_postgres(documents, metadata_dict, docs_to_upload, chunk_target_length, chunk_min_length, keep_batches): + indexer = DocsiteIndexer( + chunk_target_length=chunk_target_length, + chunk_min_length=chunk_min_length, + keep_batches=keep_batches, + ) + + conn = get_db_connection() + try: + run_migrations(conn) + register_vector_type(conn) + + batch_id = None + try: + batch_id = indexer.start_batch(conn, docs_to_upload) + chunk_count = indexer.insert_documents(conn, batch_id, documents, metadata_dict) + copied = indexer.copy_forward_missing_docs_types(conn, batch_id, docs_to_upload) + indexer.build_index(conn, batch_id) + indexer.promote_batch(conn, batch_id, chunk_count + copied) + except Exception: + if batch_id is not None: + _mark_failed(indexer, conn, batch_id) + raise + + # The new batch is already promoted and visible to readers. Pruning + # only touches older, unrelated batches. + pruned = _prune_old_batches(indexer, conn) + + return { + "target": "postgres", + "batch_id": batch_id, + "docs_types": docs_to_upload, + "chunk_count": chunk_count, + "copied_forward": copied, + "pruned_batches": pruned, + "promoted": True, + } + finally: + conn.close() - # Upload with metadata - idx = docsite_indexer.insert_documents(documents, metadata_dict) if __name__ == "__main__": - main() \ No newline at end of file + main({}) diff --git a/services/embed_docsite/pinecone_legacy_indexer.py b/services/embed_docsite/pinecone_legacy_indexer.py new file mode 100644 index 00000000..d0d2ae7c --- /dev/null +++ b/services/embed_docsite/pinecone_legacy_indexer.py @@ -0,0 +1,183 @@ +import os +import time +from datetime import datetime + +import pandas as pd +from langchain_community.document_loaders import DataFrameLoader +from langchain_openai import OpenAIEmbeddings +from langchain_pinecone import PineconeVectorStore +from pinecone import Pinecone, ServerlessSpec +from util import create_logger + +logger = create_logger("LegacyPineconeDocsiteIndexer") + +class LegacyPineconeDocsiteIndexer: + """ + Legacy Pinecone-backed docsite indexer, preserved as the write-side rollback + path for the Postgres migration. Selected via embed_docsite's `target` + payload param. Deleted alongside pinecone_legacy_search.py in the cleanup PR. + + Initialize vectorstore and insert new documents. Create a new index if needed. + + :param collection_name: Vectorstore collection name (namespace) to store documents + :param index_name: Vectorstore index name (default: docsite) + :param embeddings: LangChain embedding type (default: OpenAIEmbeddings()) + :param dimension: Embedding dimension (default: 1536 for OpenAI Embeddings) + :param max_total_collections: Max total collections in index. Delete old collections by date if exceeded after a new upload (default: 50) + """ + def __init__(self, collection_name=None, index_name="docsite", embeddings=None, dimension=1536, max_total_collections=50): + self.collection_name = collection_name if collection_name is not None else f"docsite-{datetime.now().strftime('%Y%m%d%H%M')}" + self.index_name = index_name + self.embeddings = embeddings if embeddings is not None else OpenAIEmbeddings() + self.dimension = dimension + self.max_total_collections = max_total_collections + self.pc = Pinecone(api_key=os.environ.get("PINECONE_API_KEY")) + + if not self.index_exists(): + self.create_index() + + self.index = self.pc.Index(self.index_name) + self.vectorstore = PineconeVectorStore(index_name=index_name, namespace=self.collection_name, embedding=self.embeddings) + + def insert_documents(self, inputs, metadata_dict): + """ + Create the index if it does not exist and insert the input documents. + + :param inputs: Dictionary containing name, docs_type, and doc_chunk + :param metadata_dict: Metadata dict with document titles as keys (from DocsiteProcessor) + :return: Initialized indices + """ + + # Get vector count before insertion for verification + try: + stats = self.index.describe_index_stats() + vectors_before = stats.namespaces.get(self.collection_name, {}).get("vector_count", 0) + logger.info(f"Current vector count in namespace '{self.collection_name}': {vectors_before}") + except Exception as e: + logger.warning(f"Could not get vector count before insertion: {str(e)}") + vectors_before = 0 + + df = self.preprocess_metadata(inputs=inputs, metadata_dict=metadata_dict) + logger.info(f"Input metadata preprocessed") + loader = DataFrameLoader(df, page_content_column="text") + docs = loader.load() + logger.info(f"Inputs processed into LangChain docs") + logger.info(f"Uploading {len(docs)} documents to index...") + + idx = self.vectorstore.add_documents( + documents=docs + ) + sleep_time = 10 + max_wait_time = 150 + elapsed_time = 0 + logger.info(f"Waiting up to {max_wait_time}s to verify upload count") + + while elapsed_time < max_wait_time: + time.sleep(sleep_time) + elapsed_time += sleep_time + + # Verify the upload by checking the vector count + try: + stats = self.index.describe_index_stats() + vectors_after = stats.namespaces.get(self.collection_name, {}).get("vector_count", 0) + logger.info(f"Vector count after {elapsed_time}s: {vectors_after}") + + if vectors_after >= vectors_before + len(docs): + logger.info(f"Successfully added {vectors_after - vectors_before} vectors to namespace '{self.collection_name}'") + break + else: + logger.warning(f"No new vectors were added to namespace '{self.collection_name}' after {sleep_time}s") + except Exception as e: + logger.warning(f"Could not verify vector insertion: {str(e)}") + + if vectors_after <= vectors_before: + logger.warning(f"Could not verify full dataset upload to namespace '{self.collection_name}' after {max_wait_time}s") + + self.delete_old_collections(self.max_total_collections) + + return idx + + def delete_collection(self): + """ + Deletes the entire collection (namespace) and all its contents. + This operation cannot be undone and removes both the collection structure and all vectors/documents within it. + """ + self.index.delete(delete_all=True, namespace=self.collection_name) + + def delete_old_collections(self, max_total_collections): + """Retrieve docsite uploads by collection name from Pinecone and delete them if there are more than max_total_collections.""" + + logger.info(f"Fetching outdated docsite collections") + pc = Pinecone(api_key=os.environ.get("PINECONE_API_KEY")) + index = pc.Index("docsite") + index_stats = index.describe_index_stats() + namespaces = index_stats.get('namespaces', {}).keys() + valid_namespaces = sorted( + (ns for ns in namespaces if ns.startswith("docsite-") and ns[8:].isdigit() and len(ns) in (16, 20)), + reverse=False + ) + if len(valid_namespaces) > max_total_collections: + logger.info(f"Deleting outdated docsite collections") + for old_collection in valid_namespaces[:max_total_collections]: + self.index.delete(delete_all=True, namespace=old_collection) + logger.info(f"Deleted collection {old_collection}") + + if not valid_namespaces: + logger.info(f"No valid namespaces found in the index when deleting old collections.") + + def create_index(self): + """Creates a new Pinecone index if it does not exist.""" + + if not self.index_exists(): + self.pc.create_index( + name=self.index_name, + dimension=self.dimension, + metric="cosine", + spec=ServerlessSpec(cloud="aws", region="us-east-1") + ) + while not self.pc.describe_index(self.index_name).status["ready"]: + time.sleep(1) + + def index_exists(self): + """Check if the index exists in Pinecone.""" + existing_indexes = [index_info["name"] for index_info in self.pc.list_indexes()] + + return self.index_name in existing_indexes + + def preprocess_metadata(self, inputs, page_content_column="text", add_chunk_as_metadata=False, metadata_cols=None, metadata_dict=None): + """ + Create a DataFrame for indexing from input documents and metadata. + + :param inputs: Dictionary containing name, docs_type, and doc_chunk + :param page_content_column: Name of the field which will be embedded (default: text) + :param add_chunk_as_metadata: Copy the text to embed as a separate metadata field (default: False) + :param metadata_cols: Optional list of metadata columns to include (default: None) + :param metadata_dict: Dictionary mapping names to metadata dictionaries (default: None) + :return: pandas.DataFrame with text and metadata columns + """ + + # Create DataFrame from the inputs (doc_chunk, name, docs_type) + df = pd.DataFrame(inputs) + + # Rename some columns for metadata upload + df = df.rename(columns={"doc_chunk": page_content_column, "name": "doc_title"}) + + df["doc_title"] = df["doc_title"].str.replace(".md$", "", regex=True) + + # Optionally add chunk to metadata for keyword searching + if add_chunk_as_metadata: + df["embedding_text"] = df[page_content_column] + + # Add further metadata columns if specified + if metadata_cols: + for col in metadata_cols: + df[col] = metadata_dict.get(inputs["name"], {}).get(col) + + return df + + + + + + + \ No newline at end of file diff --git a/services/embed_docsite/tests/__init__.py b/services/embed_docsite/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/embed_docsite/tests/integration/__init__.py b/services/embed_docsite/tests/integration/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/embed_docsite/tests/integration/conftest.py b/services/embed_docsite/tests/integration/conftest.py new file mode 100644 index 00000000..a9551513 --- /dev/null +++ b/services/embed_docsite/tests/integration/conftest.py @@ -0,0 +1,37 @@ +"""Fixtures for the Postgres docsite integration suite. + +Requires POSTGRES_TEST_URL (not POSTGRES_URL) pointing at a database with the pgvector extension +available and a role permitted to CREATE EXTENSION: + + docker run -d --name apollo-pgvector-test -e POSTGRES_PASSWORD=postgres \ + -p 5433:5432 pgvector/pgvector:pg16 + export POSTGRES_TEST_URL=postgresql://postgres:postgres@127.0.0.1:5433/postgres + +The repo-root conftest blocks psycopg2.connect for `unit` tests only, so this +tier connects normally. +""" + +import psycopg2 +import pytest +from embed_docsite.tests.integration.helpers import TEST_URL + + +@pytest.fixture +def clean_db(monkeypatch): + """An empty database: no docsite tables, no pgvector extension. + + Dropping the extension too means the reader's 'no extension' path is + reachable, and migrations have to prove they can recreate it. + """ + if not TEST_URL: + pytest.skip("POSTGRES_TEST_URL not set") + + conn = psycopg2.connect(TEST_URL) + conn.autocommit = True + with conn.cursor() as cur: + cur.execute("DROP TABLE IF EXISTS docsite_chunks, docsite_batches, _migrations_docs CASCADE") + cur.execute("DROP EXTENSION IF EXISTS vector CASCADE") + conn.close() + + monkeypatch.setenv("POSTGRES_URL", TEST_URL) + return TEST_URL diff --git a/services/embed_docsite/tests/integration/helpers.py b/services/embed_docsite/tests/integration/helpers.py new file mode 100644 index 00000000..0dae25fd --- /dev/null +++ b/services/embed_docsite/tests/integration/helpers.py @@ -0,0 +1,50 @@ +"""Shared helpers for the Postgres docsite integration suite.""" + +import hashlib +import os +import random + +import psycopg2 + +# Importing embed_docsite pulls in LegacyPineconeDocsiteIndexer, whose +# OpenAIEmbeddings() default arg validates credentials at construction. Dummy +# placeholders only — this suite makes no OpenAI call. +os.environ.setdefault("OPENAI_API_KEY", "sk-test-dummy") +os.environ.setdefault("PINECONE_API_KEY", "pc-test-dummy") + +TEST_URL = os.environ.get("POSTGRES_TEST_URL") + +EMBEDDING_DIMENSIONS = 1536 + + +class StubEmbeddings: + """Deterministic stand-in for OpenAIEmbeddings. + + Identical text yields an identical vector, so querying a chunk's exact text + puts that chunk at cosine distance 0 and therefore rank 1. + """ + + model = "stub-embedding-model" + + def embed_documents(self, texts): + return [self._vector(text) for text in texts] + + def embed_query(self, text): + return self._vector(text) + + @staticmethod + def _vector(text): + seed = int.from_bytes(hashlib.sha256(text.encode("utf-8")).digest()[:8], "big") + rng = random.Random(seed) + return [rng.uniform(-1.0, 1.0) for _ in range(EMBEDDING_DIMENSIONS)] + + +def query(sql, params=None): + """Run a read against the test database on its own connection.""" + conn = psycopg2.connect(TEST_URL) + try: + with conn.cursor() as cur: + cur.execute(sql, params) + return cur.fetchall() + finally: + conn.close() diff --git a/services/embed_docsite/tests/integration/test_postgres_docsite_roundtrip.py b/services/embed_docsite/tests/integration/test_postgres_docsite_roundtrip.py new file mode 100644 index 00000000..23d67992 --- /dev/null +++ b/services/embed_docsite/tests/integration/test_postgres_docsite_roundtrip.py @@ -0,0 +1,87 @@ +"""End-to-end index-then-search against a real Postgres with pgvector. +""" + +from unittest.mock import patch + +import embed_docsite.docsite_indexer as indexer_module +import pytest +from embed_docsite.embed_docsite import _upload_to_postgres +from embed_docsite.tests.integration.helpers import StubEmbeddings, query +from search_docsite.search_docsite import DocsiteSearch +from util import ApolloError + +DOCS = [ + { + "name": "webhook-guide.md", + "docs_type": "general_docs", + "doc_chunk": "Configure a webhook trigger to start a workflow when data arrives.", + }, + { + "name": "cron-guide.md", + "docs_type": "general_docs", + "doc_chunk": "Use a cron trigger to run a workflow on a fixed schedule.", + }, +] + + +def index_docs(keep_batches=2): + """Run the real Postgres upload path with stubbed embeddings.""" + with patch.object(indexer_module, "OpenAIEmbeddings", StubEmbeddings): + return _upload_to_postgres( + documents=[dict(doc) for doc in DOCS], + metadata_dict={}, + docs_to_upload=["general_docs"], + chunk_target_length=1000, + chunk_min_length=700, + keep_batches=keep_batches, + ) + + +def make_search(**kwargs): + search = DocsiteSearch(**kwargs) + search._embeddings = StubEmbeddings() + return search + + +def test_fresh_database_migrates_indexes_and_promotes(clean_db): + """The first run on an empty database used to strand in 'building', + because copy_forward's SELECT left a transaction open.""" + result = index_docs() + + assert result["promoted"] is True + rows = query("SELECT status FROM docsite_batches WHERE id = %s", (result["batch_id"],)) + assert rows[0][0] == "complete" + + +@pytest.mark.parametrize("strategy", ["semantic", "keyword", "hybrid"]) +def test_search_returns_the_indexed_chunk(clean_db, strategy): + """Semantic and hybrid used to fail with + `operator does not exist: vector <=> numeric[]` on every query.""" + index_docs() + target = DOCS[0]["doc_chunk"] + + results = make_search().search(target, strategy=strategy, top_k=3) + + assert any(result.text == target for result in results) + + +def test_reindexing_prunes_the_previous_batch(clean_db): + """Pruning never ran, so every re-index permanently added a full + docsite copy and another HNSW index.""" + first = index_docs(keep_batches=1) + second = index_docs(keep_batches=1) + + assert first["batch_id"] in second["pruned_batches"] + assert query("SELECT count(*) FROM docsite_batches WHERE id = %s", (first["batch_id"],))[0][0] == 0 + index_name = f"idx_docsite_chunks_hnsw_{first['batch_id']}" + assert query("SELECT count(*) FROM pg_class WHERE relname = %s", (index_name,))[0][0] == 0 + + +def test_reader_without_a_schema_gets_a_clear_503(clean_db): + """With migrations moved to the indexer, a reader on an un-indexed + database must explain itself rather than emit a psycopg2 traceback.""" + with pytest.raises(ApolloError) as exc: + make_search().search("anything") + + assert exc.value.code == 503 + assert "embed_docsite" in exc.value.message diff --git a/services/embed_docsite/tests/unit/__init__.py b/services/embed_docsite/tests/unit/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/services/embed_docsite/tests/unit/conftest.py b/services/embed_docsite/tests/unit/conftest.py new file mode 100644 index 00000000..1d3871be --- /dev/null +++ b/services/embed_docsite/tests/unit/conftest.py @@ -0,0 +1,14 @@ +"""Test config for embed_docsite unit tests. + +Importing `embed_docsite.pinecone_legacy_indexer` pulls in +`LegacyPineconeDocsiteIndexer`, whose module-level `OpenAIEmbeddings()` +default arg validates credentials at construction (openai 2.x / langchain-openai 1.x). +A key must therefore exist at import time. + +This test only sets dummy environment variables for unit tests. +""" + +import os + +os.environ.setdefault("OPENAI_API_KEY", "sk-test-dummy") +os.environ.setdefault("PINECONE_API_KEY", "pc-test-dummy") diff --git a/services/embed_docsite/tests/unit/fake_conn.py b/services/embed_docsite/tests/unit/fake_conn.py new file mode 100644 index 00000000..5d921d98 --- /dev/null +++ b/services/embed_docsite/tests/unit/fake_conn.py @@ -0,0 +1,107 @@ +"""A psycopg2-shaped connection that models transaction state. + +MagicMock let two defects through: `conn.autocommit = True` on a mock is an +inert attribute write, but psycopg2 raises `set_session cannot be used inside a +transaction` when a transaction is open — and a bare SELECT is enough to open +one. This models the state machine those defects depend on. +""" + +import psycopg2 +from psycopg2.extensions import STATUS_IN_TRANSACTION, STATUS_READY + + +class FakeCursor: + """Cursor serving rows from its connection's queued results.""" + + def __init__(self, conn): + self._conn = conn + self._rows = [] + self.rowcount = -1 + + def __enter__(self): + return self + + def __exit__(self, *_exc): + return False + + def execute(self, sql, params=None): + self._conn._record(sql, params) + self._rows, self.rowcount = self._conn._pop_result() + + def fetchone(self): + return self._rows[0] if self._rows else None + + def fetchall(self): + return list(self._rows) + + def close(self): + pass + + +class FakeConn: + """psycopg2-shaped connection modelling the transaction state machine. + + :param results: FIFO consumed one entry per execute(). An entry may be a + list of rows, an int meaning "no rows, but report this rowcount" (as + INSERT and DELETE do), or None. An exhausted queue yields no rows. + :param fail_on: substring; an execute() whose SQL contains it raises and + leaves the connection aborted, as psycopg2 would. + """ + + def __init__(self, results=None, fail_on=None): + # _guard off while __init__ populates state, so the autocommit check + # below does not fire before `status` exists. + object.__setattr__(self, "_guard", False) + self._results = list(results or []) + self._fail_on = fail_on + self._failed = False + self.status = STATUS_READY + self.executed = [] + self.commits = 0 + self.rollbacks = 0 + self.closed = False + self.autocommit = False + object.__setattr__(self, "_guard", True) + + def __setattr__(self, name, value): + """Only `autocommit` is guarded — it is the assignment psycopg2 rejects.""" + if name == "autocommit" and self.__dict__.get("_guard") and self.status != STATUS_READY: + raise psycopg2.ProgrammingError("set_session cannot be used inside a transaction") + object.__setattr__(self, name, value) + + def cursor(self): + return FakeCursor(self) + + def commit(self): + self.commits += 1 + self.status = STATUS_READY + self._failed = False + + def rollback(self): + self.rollbacks += 1 + self.status = STATUS_READY + self._failed = False + + def close(self): + self.closed = True + + def _record(self, sql, params): + if self._failed: + raise psycopg2.errors.InFailedSqlTransaction( + "current transaction is aborted, commands ignored until end of transaction block" + ) + self.executed.append((sql, params)) + if self._fail_on and self._fail_on in sql: + self._failed = True + self.status = STATUS_IN_TRANSACTION + raise psycopg2.ProgrammingError(f"fake failure on: {self._fail_on}") + if not self.autocommit: + self.status = STATUS_IN_TRANSACTION + + def _pop_result(self): + """Returns (rows, rowcount) for the next queued entry.""" + entry = self._results.pop(0) if self._results else None + if isinstance(entry, int): + return [], entry + rows = list(entry or []) + return rows, len(rows) diff --git a/services/embed_docsite/tests/unit/test_db_migrations.py b/services/embed_docsite/tests/unit/test_db_migrations.py new file mode 100644 index 00000000..682a78a8 --- /dev/null +++ b/services/embed_docsite/tests/unit/test_db_migrations.py @@ -0,0 +1,122 @@ +"""Unit tests for the Python migration runner. + +Mirrors the TypeScript runner in platform/src/db/migrate.ts using lexical ordering, +already-applied files skipped and an advisory lock taken. The connection and cursor +are MagicMocks. +""" + +from unittest.mock import MagicMock, patch + +import db_migrations as m + + +def make_conn(): + conn = MagicMock() + cur = MagicMock() + conn.cursor.return_value.__enter__.return_value = cur + return conn, cur + + +def test_run_migrations_takes_advisory_lock_before_applying(): + conn, cur = make_conn() + cur.fetchall.return_value = [] + + with patch.object(m, "_migration_files", return_value=[]): + m.run_migrations(conn) + + first_sql = cur.execute.call_args_list[0][0][0] + assert "pg_advisory_xact_lock" in first_sql + assert "8314" in str(cur.execute.call_args_list[0]) + + +def test_run_migrations_creates_tracking_table(): + conn, cur = make_conn() + cur.fetchall.return_value = [] + + with patch.object(m, "_migration_files", return_value=[]): + m.run_migrations(conn) + + all_sql = " ".join(str(call) for call in cur.execute.call_args_list) + assert "_migrations_docs" in all_sql + + +def test_migration_files_returns_sql_files_in_lexical_order(tmp_path): + """The sort lives in _migration_files, so it must be tested against the real + filesystem. Files are created out of order and a non-.sql file is included + to prove it is filtered.""" + (tmp_path / "0002_second.sql").write_text("SELECT 2;", encoding="utf-8") + (tmp_path / "0010_tenth.sql").write_text("SELECT 10;", encoding="utf-8") + (tmp_path / "0001_first.sql").write_text("SELECT 1;", encoding="utf-8") + (tmp_path / "notes.md").write_text("not a migration", encoding="utf-8") + + with patch.object(m, "MIGRATIONS_DIR", tmp_path): + names = [p.name for p in m._migration_files()] + + assert names == ["0001_first.sql", "0002_second.sql", "0010_tenth.sql"] + + +def test_migration_files_returns_empty_when_dir_missing(tmp_path): + with patch.object(m, "MIGRATIONS_DIR", tmp_path / "does_not_exist"): + assert m._migration_files() == [] + + +def test_run_migrations_applies_pending_files_in_the_order_given(tmp_path): + conn, cur = make_conn() + cur.fetchall.return_value = [] + + first = tmp_path / "0001_first.sql" + first.write_text("SELECT 1;", encoding="utf-8") + second = tmp_path / "0002_second.sql" + second.write_text("SELECT 2;", encoding="utf-8") + + with patch.object(m, "_migration_files", return_value=[first, second]): + applied = m.run_migrations(conn) + + assert applied == 2 + executed = [str(call) for call in cur.execute.call_args_list] + first_idx = next(i for i, e in enumerate(executed) if "SELECT 1;" in e) + second_idx = next(i for i, e in enumerate(executed) if "SELECT 2;" in e) + assert first_idx < second_idx + + +def test_run_migrations_skips_already_applied_files(tmp_path): + conn, cur = make_conn() + cur.fetchall.return_value = [("0001_first.sql",)] + + first = tmp_path / "0001_first.sql" + first.write_text("SELECT 1;", encoding="utf-8") + + with patch.object(m, "_migration_files", return_value=[first]): + applied = m.run_migrations(conn) + + assert applied == 0 + executed = " ".join(str(call) for call in cur.execute.call_args_list) + assert "SELECT 1;" not in executed + + +def test_run_migrations_commits_once_at_the_end(tmp_path): + conn, cur = make_conn() + cur.fetchall.return_value = [] + + first = tmp_path / "0001_first.sql" + first.write_text("SELECT 1;", encoding="utf-8") + + with patch.object(m, "_migration_files", return_value=[first]): + m.run_migrations(conn) + + conn.commit.assert_called_once() + + +def test_get_db_connection_does_not_run_migrations(): + """Migrations belong to the indexer, not to every reader. CREATE EXTENSION + needs privileges managed Postgres withholds, so a reader that triggers it + 500s on a deployment that never enabled the Postgres docsite backend.""" + import util + + with patch.object(util, "psycopg2") as mock_psycopg2, \ + patch.object(m, "run_migrations") as mock_run, \ + patch.dict("os.environ", {"POSTGRES_URL": "postgresql://user@host/db"}): + conn = util.get_db_connection() + + assert conn is mock_psycopg2.connect.return_value + mock_run.assert_not_called() diff --git a/services/embed_docsite/tests/unit/test_docsite_indexer.py b/services/embed_docsite/tests/unit/test_docsite_indexer.py new file mode 100644 index 00000000..76ca2400 --- /dev/null +++ b/services/embed_docsite/tests/unit/test_docsite_indexer.py @@ -0,0 +1,158 @@ +"""Unit tests for the Postgres batch-lifecycle write path. + +Every DB call goes through an explicit `conn` parameter (never looked up +internally), so these tests drive a FakeConn that models psycopg2's transaction +state. +""" + +from unittest.mock import MagicMock, patch + +import embed_docsite.docsite_indexer as m +import psycopg2 +import pytest +from embed_docsite.tests.unit.fake_conn import FakeConn +from pgvector import Vector + + +def test_register_vector_type_calls_pgvector_register(): + """Asserts a call, not transaction behaviour, so a MagicMock is right here.""" + conn = MagicMock() + with patch.object(m, "register_vector") as mock_register: + m.register_vector_type(conn) + mock_register.assert_called_once_with(conn) + + +def make_indexer(): + indexer = m.DocsiteIndexer(chunk_target_length=1000, chunk_min_length=700, keep_batches=2) + indexer._embeddings = MagicMock(model="fake-embedding-model") + return indexer + + +def test_start_batch_inserts_row_and_returns_id(): + conn = FakeConn(results=[[(7,)]]) + indexer = make_indexer() + + batch_id = indexer.start_batch(conn, ["general_docs"]) + + assert batch_id == 7 + assert conn.executed[0][1] == (["general_docs"], 1000, 700, "fake-embedding-model") + assert conn.commits == 1 + + +def test_insert_documents_embeds_and_bulk_inserts(): + conn = FakeConn() + indexer = make_indexer() + indexer._embeddings.embed_documents.return_value = [[0.1, 0.2], [0.3, 0.4]] + documents = [ + {"name": "doc-a.md", "docs_type": "general_docs", "doc_chunk": "chunk one"}, + {"name": "doc-a.md", "docs_type": "general_docs", "doc_chunk": "chunk two"}, + ] + + with patch.object(m, "execute_values") as mock_execute_values: + count = indexer.insert_documents(conn, batch_id=7, documents=documents, metadata_dict={}) + + assert count == 2 + indexer._embeddings.embed_documents.assert_called_once_with(["chunk one", "chunk two"]) + rows = mock_execute_values.call_args[0][2] + assert rows[0] == (7, "doc-a", "general_docs", 0, "chunk one", Vector([0.1, 0.2])) + assert rows[1] == (7, "doc-a", "general_docs", 1, "chunk two", Vector([0.3, 0.4])) + assert conn.commits == 1 + + +def test_insert_documents_returns_zero_for_empty_input(): + conn = FakeConn() + indexer = make_indexer() + + count = indexer.insert_documents(conn, batch_id=7, documents=[], metadata_dict={}) + + assert count == 0 + assert conn.executed == [] + indexer._embeddings.embed_documents.assert_not_called() + + +def test_copy_forward_missing_docs_types_no_op_when_all_types_present(): + conn = FakeConn() + indexer = make_indexer() + + copied = indexer.copy_forward_missing_docs_types(conn, batch_id=7, docs_types_present=m.ALL_DOCS_TYPES) + + assert copied == 0 + assert conn.executed == [] + + +def test_copy_forward_missing_docs_types_copies_from_previous_batch(): + # Queue: the SELECT finds batch 3, then the INSERT ... SELECT reports 5 rows. + conn = FakeConn(results=[[(3,)], 5]) + indexer = make_indexer() + + copied = indexer.copy_forward_missing_docs_types(conn, batch_id=7, docs_types_present=["general_docs"]) + + assert copied == 5 + assert conn.executed[-1][1] == (7, 3, ["adaptor_docs", "adaptor_functions"]) + + +def test_copy_forward_missing_docs_types_returns_zero_when_no_previous_batch(): + conn = FakeConn(results=[None]) + indexer = make_indexer() + + copied = indexer.copy_forward_missing_docs_types(conn, batch_id=7, docs_types_present=["general_docs"]) + + assert copied == 0 + + +def test_promote_batch_updates_status_and_chunk_count(): + conn = FakeConn() + indexer = make_indexer() + + indexer.promote_batch(conn, batch_id=7, chunk_count=42) + + assert conn.executed[0][1] == (42, 7) + assert conn.commits == 1 + + +def test_prune_old_batches_deletes_batches_beyond_keep_count(): + conn = FakeConn(results=[[(3,), (2,)]]) + indexer = make_indexer() + + pruned = indexer.prune_old_batches(conn, keep_batches=2) + + assert pruned == [3, 2] + assert conn.executed[0][1] == (2,) + + +def test_build_index_succeeds_when_a_prior_read_left_a_transaction_open(): + conn = FakeConn(results=[None]) + indexer = make_indexer() + indexer.copy_forward_missing_docs_types(conn, batch_id=7, docs_types_present=["general_docs"]) + + indexer.build_index(conn, batch_id=7) + + assert any("CREATE INDEX CONCURRENTLY" in sql for sql, _ in conn.executed) + assert conn.autocommit is False + + +def test_prune_old_batches_completes_after_its_own_select(): + conn = FakeConn(results=[[(3,), (2,)]]) + indexer = make_indexer() + + pruned = indexer.prune_old_batches(conn, keep_batches=2) + + assert pruned == [3, 2] + dropped = [sql for sql, _ in conn.executed if "DROP INDEX CONCURRENTLY" in sql] + assert len(dropped) == 2 + assert conn.autocommit is False + + +def test_fail_batch_recovers_an_aborted_transaction(): + conn = FakeConn(fail_on="boom") + indexer = make_indexer() + with pytest.raises(psycopg2.ProgrammingError): + with conn.cursor() as cur: + cur.execute("boom") + + indexer.fail_batch(conn, batch_id=7) + + assert conn.rollbacks == 1 + sql, params = conn.executed[-1] + assert "status = 'failed'" in sql + assert params == (7,) diff --git a/services/embed_docsite/tests/unit/test_docsite_processor.py b/services/embed_docsite/tests/unit/test_docsite_processor.py new file mode 100644 index 00000000..d390df57 --- /dev/null +++ b/services/embed_docsite/tests/unit/test_docsite_processor.py @@ -0,0 +1,67 @@ +"""Unit tests for DocsiteProcessor's pure text-processing pipeline. + +No network/DB: get_docs() (which hits GitHub) is never called directly — +these tests exercise _clean_html/_split_by_headers/_split_oversized_chunks/ +_accumulate_chunks/_chunk_adaptor_docs directly against in-memory fixtures. +""" + +from embed_docsite.docsite_processor import DocsiteProcessor + + +def make_processor(**kwargs): + return DocsiteProcessor(docs_type="general_docs", docs_to_ignore=[], **kwargs) + + +def test_clean_html_converts_tags(): + p = make_processor() + result = p._clean_html("
Hello
x bold drop")
+ assert result == "Hello\n `x` **bold** drop"
+
+
+def test_split_by_headers_splits_on_markdown_headers():
+ p = make_processor()
+ text = "# Title\ncontent one\n## Subtitle\ncontent two"
+ sections = p._split_by_headers(text)
+ assert sections == ["# Title\ncontent one", "## Subtitle\ncontent two"]
+
+
+def test_split_oversized_chunks_splits_on_newlines_when_over_target():
+ p = make_processor()
+ chunk = "a" * 5 + "\n" + "b" * 5 + "\n" + "c" * 5
+ result = p._split_oversized_chunks([chunk], target_length=8)
+ assert result == ["aaaaa", "bbbbb", "ccccc"]
+
+
+def test_accumulate_chunks_merges_up_to_target_length():
+ p = make_processor()
+ splits = ["a" * 5, "b" * 5, "c" * 5]
+ result = p._accumulate_chunks(splits, target_length=12, overlap=1, min_length=8)
+ assert result == ["aaaaabbbbb", "aaaaabbbbbccccc"]
+
+
+def test_chunk_adaptor_docs_respects_custom_target_and_min_length():
+ p = make_processor(target_length=20, min_length=15, overlap=1)
+ json_data = [{"name": "doc-a.md", "docs": "# Header\n" + ("word " * 10).strip()}]
+
+ chunks, metadata_dict = p._chunk_adaptor_docs(json_data)
+
+ assert all(c["name"] == "doc-a.md" for c in chunks)
+ assert all(c["docs_type"] == "general_docs" for c in chunks)
+ assert "doc-a.md" in metadata_dict
+
+
+def test_chunk_adaptor_docs_skips_ignored_docs():
+ p = DocsiteProcessor(docs_type="general_docs", docs_to_ignore=["skip-me.md"])
+ json_data = [{"name": "skip-me.md", "docs": "content"}]
+
+ chunks, metadata_dict = p._chunk_adaptor_docs(json_data)
+
+ assert chunks == []
+ assert metadata_dict == {}
+
+
+def test_constructor_defaults_match_previous_hardcoded_values():
+ p = DocsiteProcessor(docs_type="general_docs")
+ assert p.target_length == 1000
+ assert p.min_length == 700
+ assert p.overlap == 1
diff --git a/services/embed_docsite/tests/unit/test_embed_docsite.py b/services/embed_docsite/tests/unit/test_embed_docsite.py
new file mode 100644
index 00000000..9651ff90
--- /dev/null
+++ b/services/embed_docsite/tests/unit/test_embed_docsite.py
@@ -0,0 +1,213 @@
+"""Unit tests for embed_docsite's orchestration. DocsiteProcessor/DocsiteIndexer
+and get_db_connection are all mocked."""
+
+from unittest.mock import MagicMock, patch
+
+import embed_docsite.embed_docsite as m
+
+
+def test_main_orchestrates_full_batch_lifecycle_and_returns_summary():
+ fake_conn = MagicMock()
+ fake_indexer = MagicMock()
+ fake_indexer.start_batch.return_value = 7
+ fake_indexer.insert_documents.return_value = 10
+ fake_indexer.copy_forward_missing_docs_types.return_value = 3
+ fake_indexer.prune_old_batches.return_value = [4]
+
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([{"name": "a.md", "docs_type": "general_docs", "doc_chunk": "x"}], {"a.md": {}})
+
+ with patch.object(m, "get_db_connection", return_value=fake_conn), \
+ patch.object(m, "register_vector_type"), \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "DocsiteIndexer", return_value=fake_indexer), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ result = m.main({"docs_to_upload": ["general_docs"], "target": "postgres"})
+
+ fake_indexer.start_batch.assert_called_once_with(fake_conn, ["general_docs"])
+ fake_indexer.insert_documents.assert_called_once()
+ fake_indexer.copy_forward_missing_docs_types.assert_called_once_with(fake_conn, 7, ["general_docs"])
+ fake_indexer.build_index.assert_called_once_with(fake_conn, 7)
+ fake_indexer.promote_batch.assert_called_once_with(fake_conn, 7, 13) # 10 inserted + 3 copied forward
+ fake_indexer.prune_old_batches.assert_called_once_with(fake_conn)
+ fake_conn.close.assert_called_once()
+
+ assert result == {
+ "target": "postgres",
+ "batch_id": 7,
+ "docs_types": ["general_docs"],
+ "chunk_count": 10,
+ "copied_forward": 3,
+ "pruned_batches": [4],
+ "promoted": True,
+ }
+
+def test_main_defaults_docs_to_upload_to_all_types():
+ fake_conn = MagicMock()
+ fake_indexer = MagicMock()
+ fake_indexer.start_batch.return_value = 1
+ fake_indexer.insert_documents.return_value = 0
+ fake_indexer.copy_forward_missing_docs_types.return_value = 0
+ fake_indexer.prune_old_batches.return_value = []
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([], {})
+
+ with patch.object(m, "get_db_connection", return_value=fake_conn), \
+ patch.object(m, "register_vector_type"), \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor) as mock_processor_cls, \
+ patch.object(m, "DocsiteIndexer", return_value=fake_indexer), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ m.main({"target": "postgres"})
+
+ called_docs_types = [call.kwargs["docs_type"] for call in mock_processor_cls.call_args_list]
+ assert called_docs_types == m.ALL_DOCS_TYPES
+
+
+def test_main_defaults_to_pinecone_target():
+ """Default must match main's behavior: write to Pinecone, never open Postgres."""
+ fake_indexer = MagicMock()
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([], {})
+
+ with patch.object(m, "get_db_connection") as mock_get_conn, \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "LegacyPineconeDocsiteIndexer", return_value=fake_indexer) as mock_legacy_cls, \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test", "PINECONE_API_KEY": "pc-test"}):
+ m.main({"docs_to_upload": ["general_docs"]})
+
+ mock_legacy_cls.assert_called_once()
+ mock_get_conn.assert_not_called()
+
+
+def test_main_pinecone_target_does_not_require_postgres_url():
+ """With target=pinecone the service must not touch Postgres at all, so
+ POSTGRES_URL need not be set — matching main's dependency surface."""
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([], {})
+
+ with patch.object(m, "get_db_connection") as mock_get_conn, \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "LegacyPineconeDocsiteIndexer"), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test", "PINECONE_API_KEY": "pc-test"}, clear=True):
+ m.main({})
+
+ mock_get_conn.assert_not_called()
+
+
+def test_main_rejects_unknown_target():
+ import pytest
+ from util import ApolloError
+
+ with patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ with pytest.raises(ApolloError) as exc:
+ m.main({"target": "elasticsearch"})
+
+ assert exc.value.code == 400
+
+
+def test_postgres_upload_runs_migrations_before_registering_vector_type():
+ """Order is load-bearing, not incidental: register_vector runs
+ to_regtype('vector') and raises unless CREATE EXTENSION has already run."""
+ calls = []
+ fake_conn = MagicMock()
+ fake_indexer = MagicMock()
+ fake_indexer.start_batch.return_value = 7
+ fake_indexer.insert_documents.return_value = 1
+ fake_indexer.copy_forward_missing_docs_types.return_value = 0
+ fake_indexer.prune_old_batches.return_value = []
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([], {})
+
+ with patch.object(m, "get_db_connection", return_value=fake_conn), \
+ patch.object(m, "run_migrations", side_effect=lambda _conn: calls.append("migrate")), \
+ patch.object(m, "register_vector_type", side_effect=lambda _conn: calls.append("register")), \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "DocsiteIndexer", return_value=fake_indexer), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ m.main({"docs_to_upload": ["general_docs"], "target": "postgres"})
+
+ assert calls == ["migrate", "register"]
+
+
+def test_postgres_upload_marks_batch_failed_and_reraises():
+ """Schema allows status='failed' and nothing ever set it, so an interrupted
+ index run left a 'building' row that no later run could interpret."""
+ import pytest
+
+ fake_conn = MagicMock()
+ fake_indexer = MagicMock()
+ fake_indexer.start_batch.return_value = 7
+ fake_indexer.insert_documents.side_effect = RuntimeError("embedding API down")
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([], {})
+
+ with patch.object(m, "get_db_connection", return_value=fake_conn), \
+ patch.object(m, "run_migrations"), \
+ patch.object(m, "register_vector_type"), \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "DocsiteIndexer", return_value=fake_indexer), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ with pytest.raises(RuntimeError, match="embedding API down"):
+ m.main({"docs_to_upload": ["general_docs"], "target": "postgres"})
+
+ fake_indexer.fail_batch.assert_called_once_with(fake_conn, 7)
+ fake_conn.close.assert_called_once()
+
+
+def test_prune_failure_after_promote_does_not_mark_batch_failed():
+ """Pruning cleans up OLDER, unrelated batches. If it throws after the new
+ batch was already promoted, the new batch must stay 'complete' — it must
+ not be retroactively marked 'failed', and the call must still succeed."""
+ fake_conn = MagicMock()
+ fake_indexer = MagicMock()
+ fake_indexer.start_batch.return_value = 7
+ fake_indexer.insert_documents.return_value = 10
+ fake_indexer.copy_forward_missing_docs_types.return_value = 3
+ fake_indexer.prune_old_batches.side_effect = RuntimeError("lock conflict on DROP INDEX CONCURRENTLY")
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([{"name": "a.md", "docs_type": "general_docs", "doc_chunk": "x"}], {"a.md": {}})
+
+ with patch.object(m, "get_db_connection", return_value=fake_conn), \
+ patch.object(m, "run_migrations"), \
+ patch.object(m, "register_vector_type"), \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "DocsiteIndexer", return_value=fake_indexer), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ result = m.main({"docs_to_upload": ["general_docs"], "target": "postgres"})
+
+ fake_indexer.promote_batch.assert_called_once_with(fake_conn, 7, 13)
+ fake_indexer.fail_batch.assert_not_called()
+ fake_conn.close.assert_called_once()
+
+ assert result == {
+ "target": "postgres",
+ "batch_id": 7,
+ "docs_types": ["general_docs"],
+ "chunk_count": 10,
+ "copied_forward": 3,
+ "pruned_batches": [],
+ "promoted": True,
+ }
+
+
+def test_failed_marking_never_masks_the_original_error():
+ """Bookkeeping must not become the reported failure — the operator needs to
+ see what actually broke."""
+ import pytest
+
+ fake_conn = MagicMock()
+ fake_indexer = MagicMock()
+ fake_indexer.start_batch.return_value = 7
+ fake_indexer.insert_documents.side_effect = RuntimeError("embedding API down")
+ fake_indexer.fail_batch.side_effect = RuntimeError("connection already gone")
+ fake_processor = MagicMock()
+ fake_processor.get_preprocessed_docs.return_value = ([], {})
+
+ with patch.object(m, "get_db_connection", return_value=fake_conn), \
+ patch.object(m, "run_migrations"), \
+ patch.object(m, "register_vector_type"), \
+ patch.object(m, "DocsiteProcessor", return_value=fake_processor), \
+ patch.object(m, "DocsiteIndexer", return_value=fake_indexer), \
+ patch.dict("os.environ", {"OPENAI_API_KEY": "sk-test"}):
+ with pytest.raises(RuntimeError, match="embedding API down"):
+ m.main({"docs_to_upload": ["general_docs"], "target": "postgres"})
diff --git a/services/job_chat/retrieve_docs.py b/services/job_chat/retrieve_docs.py
index 5c0e644f..33d35a89 100644
--- a/services/job_chat/retrieve_docs.py
+++ b/services/job_chat/retrieve_docs.py
@@ -1,23 +1,24 @@
-import os
import json
+import os
+
import anthropic
+import sentry_sdk
from anthropic import (
APIConnectionError,
- BadRequestError,
AuthenticationError,
- PermissionDeniedError,
+ BadRequestError,
+ InternalServerError,
NotFoundError,
- UnprocessableEntityError,
+ PermissionDeniedError,
RateLimitError,
- InternalServerError,
+ UnprocessableEntityError,
)
-import sentry_sdk
from langfuse import observe
-from util import ApolloError, create_logger
from models import resolve_model
-from search_docsite.search_docsite import DocsiteSearch
+from search_docsite.search_docsite import resolve_backend
+from util import ApolloError, create_logger
+
from .rag_config_loader import ConfigLoader
-from streaming_util import StreamManager
logger = create_logger("job_chat.retrieve_docs")
@@ -82,7 +83,7 @@ def retrieve_knowledge(content, history, code="", adaptor="", api_key=None, stre
search_results = list(set(search_results))
search_results_sections = list(set(result.metadata["doc_title"] for result in search_results))
except Exception as e:
- logger.error(f"Pinecone search failed: {e}")
+ logger.error(f"Docsite search failed: {e}")
sentry_sdk.capture_exception(e)
# Continue with empty results - chat can still work without docs
search_results = []
@@ -163,26 +164,33 @@ def generate_queries(content, client, user_context=""):
"Failed to generate search queries - invalid response from AI service",
type="INVALID_LLM_RESPONSE",
details={"response_preview": text[:200]}
- )
+ ) from e
if len(answer_parsed) >= 4:
answer_parsed = answer_parsed[:4]
return (answer_parsed, usage)
-def search_docs(search_queries, top_k, threshold):
- """Search the docsite vector store using search queries."""
- docsite_search = DocsiteSearch()
+def search_docs(search_queries, top_k, threshold=None):
+ """Search the docsite store. Both backends run semantic search, so the
+ threshold applies identically.
+
+ Set DOCSITE_SEARCH_BACKEND=postgres to use Postgres.
+
+ :param threshold: Cosine-similarity cutoff
+ """
+ searcher = resolve_backend()()
search_results = []
for q in search_queries:
- query_search_result = docsite_search.search(
- q.get("query"),
- top_k=top_k,
- threshold=threshold,
+ query_search_result = searcher.search(
+ q.get("query"),
+ top_k=top_k,
+ threshold=threshold,
+ strategy="semantic",
docs_type="general_docs"
)
search_results.extend(query_search_result)
-
+
return search_results
def format_context(adaptor, code, history):
@@ -247,10 +255,10 @@ def call_llm(model, temperature, system_prompt, user_prompt, client, output_sche
"Unable to reach the AI service for documentation search",
type="CONNECTION_ERROR",
details=details,
- )
+ ) from e
except AuthenticationError as e:
logger.error(f"Authentication error during knowledge retrieval: {e}")
- raise ApolloError(401, "Authentication failed with AI service", type="AUTH_ERROR")
+ raise ApolloError(401, "Authentication failed with AI service", type="AUTH_ERROR") from e
except RateLimitError as e:
logger.error(f"Rate limit error during knowledge retrieval: {e}")
retry_after = int(e.response.headers.get('retry-after', 60)) if hasattr(e, 'response') else 60
@@ -259,22 +267,22 @@ def call_llm(model, temperature, system_prompt, user_prompt, client, output_sche
"Rate limit exceeded for documentation search, please try again later",
type="RATE_LIMIT",
details={"retry_after": retry_after}
- )
+ ) from e
except BadRequestError as e:
logger.error(f"Bad request error during knowledge retrieval: {e}")
- raise ApolloError(400, f"Invalid request to AI service: {str(e)}", type="BAD_REQUEST")
+ raise ApolloError(400, f"Invalid request to AI service: {str(e)}", type="BAD_REQUEST") from e
except PermissionDeniedError as e:
logger.error(f"Permission denied error during knowledge retrieval: {e}")
- raise ApolloError(403, "Not authorized to perform this action", type="FORBIDDEN")
+ raise ApolloError(403, "Not authorized to perform this action", type="FORBIDDEN") from e
except NotFoundError as e:
logger.error(f"Not found error during knowledge retrieval: {e}")
- raise ApolloError(404, "Resource not found", type="NOT_FOUND")
+ raise ApolloError(404, "Resource not found", type="NOT_FOUND") from e
except UnprocessableEntityError as e:
logger.error(f"Unprocessable entity error during knowledge retrieval: {e}")
- raise ApolloError(422, str(e), type="INVALID_REQUEST")
+ raise ApolloError(422, str(e), type="INVALID_REQUEST") from e
except InternalServerError as e:
logger.error(f"Internal server error from AI service during knowledge retrieval: {e}")
- raise ApolloError(500, "The AI service encountered an error", type="PROVIDER_ERROR")
+ raise ApolloError(500, "The AI service encountered an error", type="PROVIDER_ERROR") from e
except Exception as e:
logger.error(f"Unexpected error during LLM call for knowledge retrieval: {str(e)}")
- raise ApolloError(500, f"Unexpected error during documentation search: {str(e)}", type="UNKNOWN_ERROR")
\ No newline at end of file
+ raise ApolloError(500, f"Unexpected error during documentation search: {str(e)}", type="UNKNOWN_ERROR") from e
\ No newline at end of file
diff --git a/services/job_chat/tests/integration/test_adaptor_docs_pipeline.py b/services/job_chat/tests/integration/test_adaptor_docs_pipeline.py
index 78ec9979..f926316e 100644
--- a/services/job_chat/tests/integration/test_adaptor_docs_pipeline.py
+++ b/services/job_chat/tests/integration/test_adaptor_docs_pipeline.py
@@ -3,7 +3,6 @@
import psycopg2
import pytest
from dotenv import load_dotenv
-
from job_chat.prompt import generate_system_message
load_dotenv()
@@ -147,7 +146,7 @@ def test_generate_queries_returns_valid_structure():
print("==================TEST==================")
print("Description: Testing generate_queries returns valid JSON structure")
- from job_chat.retrieve_docs import generate_queries, get_client, format_context
+ from job_chat.retrieve_docs import format_context, generate_queries, get_client
# Step 1: Prepare test inputs
print("\n1. Preparing test inputs...")
@@ -159,7 +158,7 @@ def test_generate_queries_returns_valid_structure():
# Step 2: Call generate_queries
print("\n2. Calling generate_queries...")
client = get_client()
- queries, usage = generate_queries(question, client, user_context)
+ queries, _usage = generate_queries(question, client, user_context)
print(f" Generated {len(queries)} queries")
# Step 3: Validate structure
diff --git a/services/job_chat/tests/unit/conftest.py b/services/job_chat/tests/unit/conftest.py
index 1816dd69..ce4637da 100644
--- a/services/job_chat/tests/unit/conftest.py
+++ b/services/job_chat/tests/unit/conftest.py
@@ -4,8 +4,7 @@
module-level `OpenAIEmbeddings()` default arg validates credentials at construction
(openai 2.x / langchain-openai 1.x). A key must therefore exist at import time.
-Dummy placeholders only — unit tests mock every network seam, and the repo-root
-conftest blocks real client construction. `setdefault` lets a real key win.
+Dummy placeholders only.
"""
import os
diff --git a/services/job_chat/tests/unit/test_retrieve_docs.py b/services/job_chat/tests/unit/test_retrieve_docs.py
index 54fb3558..45b57ddf 100644
--- a/services/job_chat/tests/unit/test_retrieve_docs.py
+++ b/services/job_chat/tests/unit/test_retrieve_docs.py
@@ -14,11 +14,11 @@
from unittest.mock import MagicMock, patch
import pytest
-
+from embeddings.embeddings import SearchResult
from job_chat import retrieve_docs as rd
+from job_chat.retrieve_docs import search_docs
from util import ApolloError
-
# --- generate_queries ----------------------------------------------------------
def test_generate_queries_truncates_to_four():
@@ -76,3 +76,54 @@ def test_call_llm_wraps_unexpected_error_as_apollo_error():
assert exc.value.code == 500
assert exc.value.type == "UNKNOWN_ERROR"
+
+
+# --- search_docs ----------------------------------------------------------------
+
+
+def _fake_result(title):
+ return SearchResult(f"text for {title}", {"doc_title": title, "docs_type": "general_docs"}, 0.9)
+
+
+def test_search_docs_forwards_query_args_to_resolved_backend(monkeypatch):
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+
+ with patch.object(rd, "resolve_backend") as mock_resolve:
+ backend = mock_resolve.return_value.return_value
+ backend.search.return_value = [_fake_result("A")]
+ results = search_docs([{"query": "q"}], top_k=5)
+
+ assert [r.metadata["doc_title"] for r in results] == ["A"]
+ backend.search.assert_called_once_with(
+ "q", top_k=5, threshold=None, strategy="semantic", docs_type="general_docs"
+ )
+
+
+def test_search_docs_passes_threshold_through(monkeypatch):
+ """Threshold is a cosine-similarity cutoff that must reach the backend — this
+ regressed once already when the backend flag was introduced, and rag.yaml's
+ threshold silently stopped applying."""
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+
+ with patch.object(rd, "resolve_backend") as mock_resolve:
+ backend = mock_resolve.return_value.return_value
+ backend.search.return_value = [_fake_result("A")]
+ search_docs([{"query": "q"}], top_k=5, threshold=0.8)
+
+ backend.search.assert_called_once_with(
+ "q", top_k=5, threshold=0.8, strategy="semantic", docs_type="general_docs"
+ )
+
+
+def test_search_docs_accumulates_results_across_queries(monkeypatch):
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+
+ with patch.object(rd, "resolve_backend") as mock_resolve:
+ backend = mock_resolve.return_value.return_value
+ backend.search.side_effect = [[_fake_result("A")], [_fake_result("B")]]
+ results = search_docs([{"query": "q1"}, {"query": "q2"}], top_k=5)
+
+ assert [r.metadata["doc_title"] for r in results] == ["A", "B"]
+ assert backend.search.call_count == 2
+
+
diff --git a/services/migrations/20260728000000_docsite_batches_and_chunks.sql b/services/migrations/20260728000000_docsite_batches_and_chunks.sql
new file mode 100644
index 00000000..0b8dfa53
--- /dev/null
+++ b/services/migrations/20260728000000_docsite_batches_and_chunks.sql
@@ -0,0 +1,34 @@
+-- Docsite chunk storage (Postgres + pgvector).
+-- Applied by services/db_migrations.py; recorded in _migrations_docs.
+
+CREATE EXTENSION IF NOT EXISTS vector;
+
+CREATE TABLE IF NOT EXISTS docsite_batches (
+ id BIGSERIAL PRIMARY KEY,
+ status VARCHAR(20) NOT NULL DEFAULT 'building'
+ CHECK (status IN ('building', 'complete', 'failed')),
+ docs_types TEXT[] NOT NULL,
+ chunk_target_length INT NOT NULL,
+ chunk_min_length INT NOT NULL,
+ embedding_model VARCHAR(100) NOT NULL,
+ chunk_count INT,
+ started_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
+
+CREATE TABLE IF NOT EXISTS docsite_chunks (
+ id BIGSERIAL PRIMARY KEY,
+ batch_id BIGINT NOT NULL REFERENCES docsite_batches(id) ON DELETE CASCADE,
+ doc_title VARCHAR(500) NOT NULL,
+ docs_type VARCHAR(50) NOT NULL,
+ chunk_index INT NOT NULL,
+ text TEXT NOT NULL,
+ embedding vector(1536) NOT NULL,
+ text_search tsvector GENERATED ALWAYS AS (to_tsvector('english', text)) STORED,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE INDEX IF NOT EXISTS idx_docsite_chunks_batch ON docsite_chunks(batch_id);
+CREATE INDEX IF NOT EXISTS idx_docsite_chunks_doc_title ON docsite_chunks(batch_id, doc_title);
+CREATE INDEX IF NOT EXISTS idx_docsite_chunks_docs_type ON docsite_chunks(batch_id, docs_type);
+CREATE INDEX IF NOT EXISTS idx_docsite_chunks_fts ON docsite_chunks USING gin(text_search);
diff --git a/services/search_docsite/README.md b/services/search_docsite/README.md
index ecbe0146..c7380344 100644
--- a/services/search_docsite/README.md
+++ b/services/search_docsite/README.md
@@ -26,16 +26,24 @@ bun py search_docsite tmp/payload.json -O
## Implementation
The service uses the DocsiteSearch class to query the database (Pinecone). It embeds semantic search queries using OpenAI.
+To compare backends on the same query, run the service twice with different
+`backend` values and diff the results. This replaces the shadow-mode comparison
+that was considered for the Postgres migration.
+
## Payload Reference
The input payload is a JSON object with the following structure:
```js
{
- "query": "What is Asana", // Input query
- "collection_name": "Docsite-20250225", // Name of the collection in the vector database
- "docs_type": "adaptor_docs", // Filter for document type adaptor_docs, adaptor_functions, general_docs (optional)
- "doc_title": "Asana", // Filter for document title (optional)
- "top_k": 5 // Adjust the number of search results (optional)
+ "query": "What is Asana", // Input query (required)
+ "backend": "pinecone", // 'pinecone' | 'postgres'. Defaults to DOCSITE_SEARCH_BACKEND, itself defaulting to pinecone.
+ "docs_type": "adaptor_docs", // Filter for adaptor_docs | adaptor_functions | general_docs (optional)
+ "doc_title": "Asana", // Filter for document title (optional)
+ "top_k": 5, // Number of search results (optional)
+ "threshold": 0.8, // Cosine cutoff. Only valid with strategy 'semantic'. (optional)
+ "strategy": "semantic", // Postgres backend only: 'semantic' | 'keyword' | 'hybrid'
+ "batch_id": 12, // Postgres backend only: pin a specific batch (optional)
+ "collection_name": "docsite-..." // Pinecone backend only: pin a namespace (optional)
}
```
diff --git a/services/search_docsite/pinecone_legacy_search.py b/services/search_docsite/pinecone_legacy_search.py
new file mode 100644
index 00000000..384f7ef3
--- /dev/null
+++ b/services/search_docsite/pinecone_legacy_search.py
@@ -0,0 +1,111 @@
+import os
+
+from embeddings.embeddings import SearchResult
+from langchain_openai import OpenAIEmbeddings
+from langchain_pinecone import PineconeVectorStore
+from pinecone import Pinecone
+from util import ApolloError, create_logger
+
+logger = create_logger("LegacyPineconeDocsiteSearch")
+
+
+class LegacyPineconeDocsiteSearch:
+ """
+ Legacy Pinecone-backed docsite search, still the default backend and the
+ rollback path for the Postgres migration. Selected via resolve_backend()
+ in services/search_docsite/search_docsite.py.
+
+ :param collection_name: Vectorstore collection name (namespace) to store documents
+ :param index_name: Vectorstore index name (default: docsite)
+ :param default_top_k: Default number of results to return (default: 5)
+ :param embeddings: LangChain embedding type (default: OpenAIEmbeddings())
+ """
+ def __init__(self, collection_name=None, index_name="docsite", default_top_k=5, embeddings=None):
+ self.index_client = index_name
+ self.default_top_k = default_top_k
+ self.embeddings = embeddings if embeddings is not None else OpenAIEmbeddings()
+
+ if collection_name is None:
+ logger.info("Collection name not provided; retrieving the most recent collection name.")
+ collection_name = self._get_most_recent_namespace()
+
+ self.collection_name = collection_name
+ self.vectorstore = PineconeVectorStore(index_name=index_name, namespace=collection_name, embedding=self.embeddings)
+
+ def search(self, query, top_k=None, threshold=None, strategy='semantic', doc_title=None, docs_type=None):
+ filters = self._build_filter(doc_title=doc_title, docs_type=docs_type)
+ logger.info("Metadata filters built")
+
+ if strategy != 'semantic':
+ raise ApolloError(
+ 400,
+ f"The Pinecone backend only supports strategy='semantic', got '{strategy}'",
+ type="BAD_REQUEST",
+ )
+
+ return self._semantic_search(query=query, top_k=top_k, threshold=threshold, filters=filters)
+
+ def _semantic_search(self, query, top_k=None, threshold=None, filters=None):
+ if top_k is None and threshold is None:
+ top_k = self.default_top_k
+
+ max_k = top_k or 50
+
+ scored_docs = self.vectorstore.similarity_search_with_score(
+ query=query,
+ k=max_k,
+ filter=filters
+ )
+
+ logger.info(f"Similar documents retrieved: {len(scored_docs)}")
+
+ results = []
+ for doc, score in scored_docs:
+ if threshold is not None and score < threshold:
+ continue
+
+ if top_k is not None and len(results) >= top_k and threshold is None:
+ break
+
+ results.append(SearchResult(doc.page_content, doc.metadata, score))
+
+ logger.info(f"Filtered to {len(results)} results")
+ return results
+
+ def _build_filter(self, **kwargs):
+ conditions = []
+
+ if kwargs.get('doc_title'):
+ conditions.append({"doc_title": {"$eq": kwargs['doc_title']}})
+
+ if kwargs.get('docs_type'):
+ conditions.append({"docs_type": {"$eq": kwargs['docs_type']}})
+
+ if not conditions:
+ return None
+
+ if len(conditions) == 1:
+ return conditions[0]
+
+ return {"$and": conditions}
+
+ def _get_most_recent_namespace(self):
+ pc = Pinecone(api_key=os.environ.get("PINECONE_API_KEY"))
+ index = pc.Index("docsite")
+ index_stats = index.describe_index_stats()
+ namespaces = index_stats.get('namespaces', {}).keys()
+
+ # The indexer names namespaces docsite-%Y%m%d%H%M (20 chars); namespaces
+ # created by hand use docsite-%Y%m%d (16). Accept both, so a namespace the
+ # indexer just wrote is discoverable.
+ valid_namespaces = sorted(
+ (ns for ns in namespaces if ns.startswith("docsite-") and ns[8:].isdigit() and len(ns) in (16, 20)),
+ reverse=True
+ )
+
+ if not valid_namespaces:
+ raise ApolloError(404, "No valid namespaces found in the index", type="NOT_FOUND")
+
+ most_recent_namespace = valid_namespaces[0]
+ logger.info(f"Most recent docsite collection name found: {most_recent_namespace}")
+ return most_recent_namespace
diff --git a/services/search_docsite/search_docsite.py b/services/search_docsite/search_docsite.py
index f7355ea7..7638a55a 100644
--- a/services/search_docsite/search_docsite.py
+++ b/services/search_docsite/search_docsite.py
@@ -1,174 +1,265 @@
import os
+
+import psycopg2
from dotenv import load_dotenv
-from pinecone import Pinecone
-from langchain_pinecone import PineconeVectorStore
-from langchain_openai import OpenAIEmbeddings
-from util import create_logger, ApolloError
from embeddings.embeddings import SearchResult
+from langchain_openai import OpenAIEmbeddings
+from pgvector import Vector
+from pgvector.psycopg2 import register_vector
+from search_docsite.pinecone_legacy_search import LegacyPineconeDocsiteSearch
+from util import ApolloError, create_logger, get_db_connection
+
logger = create_logger("DocsiteSearch")
+SCHEMA_MISSING_MESSAGE = (
+ "Docsite schema not initialised — run embed_docsite with target=postgres"
+)
+
+
+def register_vector_type(conn):
+ """Register the pgvector adapter on this connection."""
+ register_vector(conn)
+
class DocsiteSearch:
"""
- Initialize the docsite vectorstore and search it with optional metadata filters.
-
- :param collection_name: Vectorstore collection name (namespace) to store documents
- :param index_name: Vectorstore index name (default: docsite)
+ Search embedded docsite chunks in Postgres using semantic (pgvector cosine),
+ keyword (Postgres full-text search), or hybrid (Reciprocal Rank Fusion) strategies.
+
+ :param batch_id: Explicit batch to search. If None, resolves to the newest 'complete' batch.
:param default_top_k: Default number of results to return (default: 5)
- :param embeddings: LangChain embedding type (default: OpenAIEmbeddings())
"""
- def __init__(self, collection_name=None, index_name="docsite", default_top_k=5, embeddings=OpenAIEmbeddings()):
- self.index_client = index_name
+
+ def __init__(self, batch_id=None, default_top_k=5):
self.default_top_k = default_top_k
+ self._explicit_batch_id = batch_id
+ self._embeddings = None
+
+ @property
+ def embeddings(self):
+ if self._embeddings is None:
+ self._embeddings = OpenAIEmbeddings()
+ return self._embeddings
- if collection_name is None:
- logger.info("Collection name not provided; retrieving the most recent collection name.")
- collection_name = self._get_most_recent_namespace()
+ def _connect(self):
+ """Open a connection with pgvector registered.
+ """
+ conn = get_db_connection()
+ try:
+ register_vector_type(conn)
+ except psycopg2.ProgrammingError as exc:
+ conn.close()
+ raise ApolloError(503, SCHEMA_MISSING_MESSAGE, type="DATABASE_ERROR") from exc
+ return conn
- self.collection_name = collection_name
- self.vectorstore = PineconeVectorStore(index_name=index_name, namespace=collection_name, embedding=embeddings)
-
def search(self, query, top_k=None, threshold=None, strategy='semantic', doc_title=None, docs_type=None):
"""
- Search database with optional filters.
+ Search docsite_chunks with optional filters.
:param query: Search query string
:param top_k: Number of results to return
- :param threshold: Score threshold for semantic search
- :param strategy: Search strategy (default: 'semantic')
+ :param threshold: Cosine-similarity cutoff. Valid only for
+ strategy='semantic'; raises for other strategies.
+ :param strategy: 'semantic' | 'keyword' | 'hybrid' (default: 'semantic')
:param doc_title: Filter by document title
:param docs_type: Filter by document type
:return: List of SearchResult objects
"""
- filters = self._build_filter(doc_title=doc_title, docs_type=docs_type)
- logger.info("Metadata filters built")
+ if threshold is not None and strategy != 'semantic':
+ raise ApolloError(
+ 400,
+ f"threshold is only supported for strategy='semantic', got '{strategy}'",
+ type="BAD_REQUEST",
+ )
+
+ conn = self._connect()
+ try:
+ batch_id = self._explicit_batch_id or self._resolve_current_batch(conn)
- if strategy == 'semantic':
- return self._semantic_search(query=query, top_k=top_k, threshold=threshold, filters=filters)
+ if strategy == 'semantic':
+ return self._semantic_search(conn, batch_id, query, top_k, threshold, doc_title, docs_type)
+ if strategy == 'keyword':
+ return self._keyword_search(conn, batch_id, query, top_k, doc_title, docs_type)
+ if strategy == 'hybrid':
+ return self._hybrid_search(conn, batch_id, query, top_k, doc_title, docs_type)
- def _semantic_search(self, query, top_k=None, threshold=None, filters=None):
- """Search the vectorstore using semantic search."""
+ raise ApolloError(400, f"Unknown search strategy: {strategy}", type="BAD_REQUEST")
+ finally:
+ conn.close()
+
+ def _resolve_current_batch(self, conn):
+ """Find the newest complete batch id."""
+ try:
+ with conn.cursor() as cur:
+ cur.execute("SELECT id FROM docsite_batches WHERE status = 'complete' ORDER BY id DESC LIMIT 1")
+ row = cur.fetchone()
+ except psycopg2.errors.UndefinedTable as exc:
+ raise ApolloError(503, SCHEMA_MISSING_MESSAGE, type="DATABASE_ERROR") from exc
+ if row is None:
+ raise ApolloError(404, "No complete docsite batch found", type="NOT_FOUND")
+ return row[0]
+
+ def _semantic_search(self, conn, batch_id, query, top_k, threshold, doc_title, docs_type):
if top_k is None and threshold is None:
top_k = self.default_top_k
-
max_k = top_k or 50
-
- scored_docs = self.vectorstore.similarity_search_with_score(
- query=query,
- k=max_k,
- filter=filters
- )
-
- logger.info(f"Similar documents retrieved: {len(scored_docs)}")
-
+
+ query_embedding = Vector(self.embeddings.embed_query(query))
+
+ sql = """
+ SELECT text, doc_title, docs_type, 1 - (embedding <=> %(query_embedding)s) AS score
+ FROM docsite_chunks
+ WHERE batch_id = %(batch_id)s
+ AND (%(doc_title)s IS NULL OR doc_title = %(doc_title)s)
+ AND (%(docs_type)s IS NULL OR docs_type = %(docs_type)s)
+ ORDER BY embedding <=> %(query_embedding)s
+ LIMIT %(max_k)s
+ """
+ params = {
+ "query_embedding": query_embedding, "batch_id": batch_id,
+ "doc_title": doc_title, "docs_type": docs_type, "max_k": max_k,
+ }
+ with conn.cursor() as cur:
+ cur.execute(sql, params)
+ rows = cur.fetchall()
+
results = []
- for doc, score in scored_docs:
+ for text, title, dtype, score in rows:
if threshold is not None and score < threshold:
continue
-
- # If we've reached top_k docs and no threshold is set, stop
if top_k is not None and len(results) >= top_k and threshold is None:
break
-
- results.append(SearchResult(doc.page_content, doc.metadata, score))
-
- logger.info(f"Filtered to {len(results)} results")
+ results.append(SearchResult(text, {"doc_title": title, "docs_type": dtype}, score))
+
+ logger.info(f"Semantic search returned {len(results)} results")
return results
-
- def _build_filter(self, **kwargs):
- """Build filter conditions to search the vectorstore."""
- conditions = []
-
- # Add exact match conditions
- if kwargs.get('doc_title'):
- conditions.append({"doc_title": {"$eq": kwargs['doc_title']}})
-
-
- if kwargs.get('docs_type'):
- conditions.append({"docs_type": {"$eq": kwargs['docs_type']}})
-
- # If no conditions were added, return None
- if not conditions:
- return None
-
- # If only one condition, return it directly
- if len(conditions) == 1:
- return conditions[0]
-
- # If multiple conditions, combine them with $and
- return {"$and": conditions}
-
- def _get_most_recent_namespace(self):
- """Retrieve the most recent docsite upload by collection name from Pinecone."""
-
- pc = Pinecone(api_key=os.environ.get("PINECONE_API_KEY"))
- index = pc.Index("docsite")
- index_stats = index.describe_index_stats()
- namespaces = index_stats.get('namespaces', {}).keys()
-
- valid_namespaces = sorted(
- (ns for ns in namespaces if ns.startswith("docsite-") and ns[8:].isdigit() and len(ns) == 16),
- reverse=True
- )
- if not valid_namespaces:
- raise ApolloError(404, "No valid namespaces found in the index.", type="NOT_FOUND")
+ def _keyword_search(self, conn, batch_id, query, top_k, doc_title, docs_type):
+ max_k = top_k or self.default_top_k
- most_recent_namespace = valid_namespaces[0]
- logger.info(f"Most recent docsite collection name found: {most_recent_namespace}")
- return most_recent_namespace
+ sql = """
+ SELECT text, doc_title, docs_type,
+ ts_rank_cd(text_search, plainto_tsquery('english', %(query)s)) AS score
+ FROM docsite_chunks
+ WHERE batch_id = %(batch_id)s
+ AND text_search @@ plainto_tsquery('english', %(query)s)
+ AND (%(doc_title)s IS NULL OR doc_title = %(doc_title)s)
+ AND (%(docs_type)s IS NULL OR docs_type = %(docs_type)s)
+ ORDER BY score DESC
+ LIMIT %(max_k)s
+ """
+ params = {"query": query, "batch_id": batch_id, "doc_title": doc_title, "docs_type": docs_type, "max_k": max_k}
+ with conn.cursor() as cur:
+ cur.execute(sql, params)
+ rows = cur.fetchall()
+
+ results = [SearchResult(text, {"doc_title": title, "docs_type": dtype}, score) for text, title, dtype, score in rows]
+ logger.info(f"Keyword search returned {len(results)} results")
+ return results
+
+ def _hybrid_search(self, conn, batch_id, query, top_k, doc_title, docs_type):
+ max_k = top_k or self.default_top_k
+ candidate_k = 50
+
+ query_embedding = Vector(self.embeddings.embed_query(query))
+
+ sql = """
+ WITH semantic AS (
+ SELECT id, text, doc_title, docs_type,
+ ROW_NUMBER() OVER (ORDER BY embedding <=> %(query_embedding)s) AS rnk
+ FROM docsite_chunks
+ WHERE batch_id = %(batch_id)s
+ AND (%(doc_title)s IS NULL OR doc_title = %(doc_title)s)
+ AND (%(docs_type)s IS NULL OR docs_type = %(docs_type)s)
+ ORDER BY embedding <=> %(query_embedding)s
+ LIMIT %(candidate_k)s
+ ),
+ keyword AS (
+ SELECT id, text, doc_title, docs_type,
+ ROW_NUMBER() OVER (ORDER BY ts_rank_cd(text_search, plainto_tsquery('english', %(query)s)) DESC) AS rnk
+ FROM docsite_chunks
+ WHERE batch_id = %(batch_id)s
+ AND text_search @@ plainto_tsquery('english', %(query)s)
+ AND (%(doc_title)s IS NULL OR doc_title = %(doc_title)s)
+ AND (%(docs_type)s IS NULL OR docs_type = %(docs_type)s)
+ ORDER BY rnk
+ LIMIT %(candidate_k)s
+ )
+ SELECT COALESCE(s.text, k.text) AS text,
+ COALESCE(s.doc_title, k.doc_title) AS doc_title,
+ COALESCE(s.docs_type, k.docs_type) AS docs_type,
+ COALESCE(1.0::float8 / (60 + s.rnk), 0) + COALESCE(1.0::float8 / (60 + k.rnk), 0) AS score
+ FROM semantic s FULL OUTER JOIN keyword k ON s.id = k.id
+ ORDER BY score DESC
+ LIMIT %(max_k)s
+ """
+ params = {
+ "query_embedding": query_embedding, "query": query, "batch_id": batch_id,
+ "doc_title": doc_title, "docs_type": docs_type, "candidate_k": candidate_k, "max_k": max_k,
+ }
+ with conn.cursor() as cur:
+ cur.execute(sql, params)
+ rows = cur.fetchall()
+
+ results = [
+ SearchResult(text, {"doc_title": title, "docs_type": dtype}, float(score))
+ for text, title, dtype, score in rows
+ ]
+ logger.info(f"Hybrid search returned {len(results)} results")
+ return results
+
+
+BACKEND_INDEX_PARAMS = {
+ "postgres": ["batch_id", "default_top_k"],
+ "pinecone": ["collection_name", "index_name", "default_top_k", "embeddings"],
+}
+
+
+def resolve_backend(override=None):
+ """Return the search class for the configured backend.
+
+ :param override: Backend name that is prioritised over DOCSITE_SEARCH_BACKEND
+ """
+ name = override or os.environ.get("DOCSITE_SEARCH_BACKEND", "pinecone")
+ if name not in BACKEND_INDEX_PARAMS:
+ raise ApolloError(400, f"Unknown backend '{name}'. Expected 'pinecone' or 'postgres'", type="BAD_REQUEST")
+ return DocsiteSearch if name == "postgres" else LegacyPineconeDocsiteSearch
def main(data):
logger.info("Starting...")
required_fields = ["query"]
-
missing = [field for field in required_fields if field not in data]
-
if missing:
logger.error(f"Missing required fields in data: {', '.join(missing)}")
- return
+ return None
- index_params = {}
- search_params = {"query": data["query"]}
+ backend = data.get("backend") or os.environ.get("DOCSITE_SEARCH_BACKEND", "pinecone")
+ search_cls = resolve_backend(backend)
- # Add optional parameters
+ search_params = {"query": data["query"]}
optional_search_params = ["docs_type", "doc_title", "top_k", "threshold", "strategy"]
- optional_index_params = ["collection_name", "index_name", "default_top_k", "embeddings"]
-
for key in optional_search_params:
if key in data:
search_params[key] = data[key]
- for key in optional_index_params:
- if key in data:
- index_params[key] = data[key]
+ index_params = {key: data[key] for key in BACKEND_INDEX_PARAMS[backend] if key in data}
- # Set API keys
load_dotenv(override=True)
- OPENAI_API_KEY = os.environ.get('OPENAI_API_KEY')
- PINECONE_API_KEY = os.environ.get('PINECONE_API_KEY')
-
- # Check for missing keys
- missing_keys = []
-
- if not OPENAI_API_KEY:
- missing_keys.append("OPENAI_API_KEY")
- if not PINECONE_API_KEY:
- missing_keys.append("PINECONE_API_KEY")
-
- if missing_keys:
- msg = f"Missing API keys: {', '.join(missing_keys)}"
+ openai_api_key = os.environ.get('OPENAI_API_KEY')
+ if not openai_api_key:
+ msg = "Missing API key: OPENAI_API_KEY"
logger.error(msg)
- raise ApolloError(500, f"Missing API keys: {', '.join(missing_keys)}", type="BAD_REQUEST")
+ raise ApolloError(500, msg, type="BAD_REQUEST")
- # Initialize search engine
- docsite_search = DocsiteSearch(**index_params)
- logger.info("Docsite database initialised")
+ logger.info(f"Searching docsite via the {backend} backend")
+
+ docsite_search = search_cls(**index_params)
results = docsite_search.search(**search_params)
-
+
return [result.to_json() for result in results]
+
if __name__ == "__main__":
- main()
\ No newline at end of file
+ main({})
diff --git a/services/search_docsite/tests/conftest.py b/services/search_docsite/tests/conftest.py
index d066197b..eb256c39 100644
--- a/services/search_docsite/tests/conftest.py
+++ b/services/search_docsite/tests/conftest.py
@@ -5,9 +5,7 @@
credentials at construction time. That happens when the test module is imported,
before any test runs — so a key must exist in the environment or import fails.
-These are dummy placeholders only: unit tests inject mocks for every network
-seam and the repo-root conftest additionally blocks real client construction, so
-no real key is ever used. `setdefault` means a real key (from services/.env) wins.
+We only set dummy environment variables for unit tests.
"""
import os
diff --git a/services/search_docsite/tests/eval/__init__.py b/services/search_docsite/tests/eval/__init__.py
new file mode 100644
index 00000000..e69de29b
diff --git a/services/search_docsite/tests/eval/golden_queries.yaml b/services/search_docsite/tests/eval/golden_queries.yaml
new file mode 100644
index 00000000..9f4cbad1
--- /dev/null
+++ b/services/search_docsite/tests/eval/golden_queries.yaml
@@ -0,0 +1,109 @@
+# Representative job_chat/general-docs queries with expected doc titles.
+#
+# recall@k counts a query as a hit when ANY expected title appears in the top k,
+# so each entry lists every doc that would be a reasonable answer, not only the
+# single best one. Titles are docsite filenames without .md, matching
+# docsite_chunks.doc_title.
+#
+# category is 'conceptual' (natural-language questions) or 'keyword' (exact-term
+# lookups). run_eval reports recall per category, because hybrid search is
+# expected to help on keyword lookups and tie on conceptual ones — a blended
+# figure would hide exactly that effect.
+#
+# Labels were derived from the general_docs corpus via full-text search on each
+# query's key terms, deliberately not from vector-search output: labelling from
+# the thing under test would be circular.
+queries:
+ # --- conceptual -------------------------------------------------------------
+ - query: "how do I configure a webhook trigger"
+ category: conceptual
+ expected_doc_titles: ["triggers", "webhook-auth"]
+ - query: "how do I set up a cron trigger"
+ category: conceptual
+ expected_doc_titles: ["triggers"]
+ - query: "what is a run in OpenFn"
+ category: conceptual
+ expected_doc_titles: ["terminology", "glossary"]
+ - query: "how do I use collections to store state between runs"
+ category: conceptual
+ expected_doc_titles: ["collections", "state", "cli-collections"]
+ - query: "how do I configure a credential for an adaptor"
+ category: conceptual
+ expected_doc_titles: ["credentials", "manage-credentials"]
+ - query: "what is the difference between a job and a workflow"
+ category: conceptual
+ expected_doc_titles: ["terminology", "glossary", "workflows"]
+ - query: "how do I deploy a project using the CLI"
+ category: conceptual
+ expected_doc_titles: ["cli-sync", "portability"]
+ - query: "how do I debug a failed run"
+ category: conceptual
+ expected_doc_titles: ["troubleshooting", "rerunning-workflow", "inspect-runs"]
+ - query: "how do I write a data transform function"
+ category: conceptual
+ expected_doc_titles: ["data-transformation", "job-writing-guide", "operations"]
+ - query: "how do I control who can access a project"
+ category: conceptual
+ expected_doc_titles: ["collaboration", "user-roles-permissions"]
+ - query: "how do I set up a sandbox environment"
+ category: conceptual
+ expected_doc_titles: ["sandboxes"]
+ - query: "how long is my run data kept"
+ category: conceptual
+ expected_doc_titles: ["retention-periods", "io-data-storage", "security-for-devs"]
+ - query: "how do I get notified when a workflow fails"
+ category: conceptual
+ expected_doc_titles: ["notifications"]
+ - query: "how do I version control my project with GitHub"
+ category: conceptual
+ expected_doc_titles: ["link-to-gh", "cli-sync"]
+ - query: "what security measures does OpenFn have"
+ category: conceptual
+ expected_doc_titles: ["security", "security-compliance"]
+
+ # --- keyword ----------------------------------------------------------------
+ - query: "state.cursor"
+ category: keyword
+ expected_doc_titles: ["using-cursors"]
+ - query: "--force flag on project push"
+ category: keyword
+ expected_doc_titles: ["cli-sync"]
+ - query: "webhook auth method"
+ category: keyword
+ expected_doc_titles: ["webhook-auth"]
+ - query: "workflow snapshots"
+ category: keyword
+ expected_doc_titles: ["workflow-snapshots"]
+ - query: "lazy state operator"
+ category: keyword
+ expected_doc_titles: ["lazy-state-operator"]
+ - query: "dataValue"
+ category: keyword
+ expected_doc_titles: ["job-snippets"]
+ - query: "each operation"
+ category: keyword
+ expected_doc_titles: ["operations"]
+ - query: "rerun a workflow"
+ category: keyword
+ expected_doc_titles: ["rerunning-workflow"]
+ - query: "collections CLI commands"
+ category: keyword
+ expected_doc_titles: ["cli-collections"]
+ - query: "activity history"
+ category: keyword
+ expected_doc_titles: ["activity-history"]
+ - query: "project.yaml"
+ category: keyword
+ expected_doc_titles: ["portability-v3"]
+ - query: "step editor"
+ category: keyword
+ expected_doc_titles: ["step-editor"]
+ - query: "git branch"
+ category: keyword
+ expected_doc_titles: ["working-with-branches", "cli-sync"]
+ - query: "sandbox"
+ category: keyword
+ expected_doc_titles: ["sandboxes"]
+ - query: "security compliance"
+ category: keyword
+ expected_doc_titles: ["security-compliance"]
diff --git a/services/search_docsite/tests/eval/run_eval.py b/services/search_docsite/tests/eval/run_eval.py
new file mode 100644
index 00000000..f20b6c10
--- /dev/null
+++ b/services/search_docsite/tests/eval/run_eval.py
@@ -0,0 +1,223 @@
+"""Offline recall@k + latency comparison between docsite search backends.
+
+Usage: poetry run python -m search_docsite.tests.eval.run_eval
+
+Compares Postgres semantic and hybrid search at each indexed chunk size against
+the legacy Pinecone baseline, over the golden query set. Recall is reported per
+query category, because hybrid is expected to help on keyword lookups and tie on
+conceptual ones — a blended figure would hide that.
+
+Requires a Postgres batch per chunk size in CHUNK_SIZES; configurations without
+one are skipped. Index them with embed_docsite, passing chunk_target_length,
+chunk_min_length, and keep_batches >= 3 so earlier batches are not pruned.
+
+Golden queries live in services/search_docsite/tests/eval/golden_queries.yaml
+— edit that file to add queries or curate expected_doc_titles.
+"""
+
+import time
+from pathlib import Path
+
+import yaml
+from dotenv import load_dotenv
+
+# Needed as LegacyPineconeDocsiteSearch evaluates OpenAIEmbeddings() as a
+# default argument.
+load_dotenv()
+
+from search_docsite.pinecone_legacy_search import LegacyPineconeDocsiteSearch
+from search_docsite.search_docsite import DocsiteSearch
+from util import get_db_connection
+
+GOLDEN_QUERIES_PATH = Path(__file__).parent / "golden_queries.yaml"
+
+# Chunk sizes to evaluate, matching docsite_batches.chunk_target_length.
+CHUNK_SIZES = [1000, 1800, 2500]
+STRATEGIES = ["semantic", "hybrid"]
+
+# The migration-fidelity pairing: same strategy and chunk size, different store.
+AGREEMENT_PAIR = ("Pinecone semantic", "Postgres semantic (1000)")
+
+
+def resolve_batch_id(chunk_target_length):
+ """Newest complete batch indexed at the given chunk size, or None if there is none.
+
+ Batch ids are environment-specific, so the eval looks them up by the indexing
+ configuration recorded on each batch rather than hardcoding them.
+ """
+ conn = get_db_connection()
+ try:
+ with conn.cursor() as cur:
+ cur.execute(
+ "SELECT id FROM docsite_batches WHERE status = 'complete' "
+ "AND chunk_target_length = %s ORDER BY id DESC LIMIT 1",
+ (chunk_target_length,),
+ )
+ row = cur.fetchone()
+ finally:
+ conn.close()
+
+ return row[0] if row else None
+
+
+def compute_recall_at_k(retrieved_titles, expected_titles):
+ """True if any expected title appears among the retrieved titles."""
+ return bool(set(retrieved_titles) & set(expected_titles))
+
+
+def compute_agreement(report_a, report_b):
+ """Doc-title overlap between two backends' result sets, per query and averaged.
+
+ Needs no ground truth, so it gives a usable comparison signal before
+ golden_queries.yaml is curated.
+ """
+ per_query = []
+ for a, b in zip(report_a["per_query"], report_b["per_query"], strict=True):
+ titles_a = {t for t in a["retrieved_titles"] if t is not None}
+ titles_b = {t for t in b["retrieved_titles"] if t is not None}
+ union = titles_a | titles_b
+ overlap = titles_a & titles_b
+ per_query.append({
+ "query": a["query"],
+ "overlap": len(overlap),
+ "union": len(union),
+ "jaccard": (len(overlap) / len(union)) if union else 0.0,
+ })
+
+ mean_jaccard = (sum(q["jaccard"] for q in per_query) / len(per_query)) if per_query else 0.0
+ return {"per_query": per_query, "mean_jaccard": mean_jaccard}
+
+
+def run_eval(golden_queries, make_backend, strategy, top_k=5):
+ """Run every golden query against one backend/strategy and return a report dict."""
+ backend = make_backend()
+ per_query = []
+ latencies = []
+ scored = 0
+ skipped = 0
+ hits = 0
+ by_category = {}
+
+ for item in golden_queries:
+ query = item["query"]
+ expected_titles = item.get("expected_doc_titles", [])
+ category = item.get("category")
+
+ start = time.time()
+ results = backend.search(query, top_k=top_k, strategy=strategy, docs_type="general_docs")
+ elapsed = time.time() - start
+ latencies.append(elapsed)
+
+ retrieved_titles = [r.metadata.get("doc_title") for r in results]
+
+ if expected_titles:
+ hit = compute_recall_at_k(retrieved_titles, expected_titles)
+ scored += 1
+ hits += int(hit)
+ if category:
+ counts = by_category.setdefault(category, {"scored": 0, "hits": 0})
+ counts["scored"] += 1
+ counts["hits"] += int(hit)
+ else:
+ hit = None
+ skipped += 1
+
+ per_query.append({
+ "query": query,
+ "category": category,
+ "retrieved_titles": retrieved_titles,
+ "expected_titles": expected_titles,
+ "hit": hit,
+ "latency_s": elapsed,
+ })
+
+ latencies_sorted = sorted(latencies)
+ p50 = latencies_sorted[len(latencies_sorted) // 2] if latencies_sorted else 0.0
+ p95_index = min(len(latencies_sorted) - 1, int(len(latencies_sorted) * 0.95)) if latencies_sorted else 0
+ p95 = latencies_sorted[p95_index] if latencies_sorted else 0.0
+
+ return {
+ "recall_at_k": (hits / scored) if scored else None,
+ "recall_by_category": {
+ name: {"recall": c["hits"] / c["scored"], "scored": c["scored"]}
+ for name, c in sorted(by_category.items())
+ },
+ "queries_scored": scored,
+ "queries_skipped": skipped,
+ "p50_latency_s": p50,
+ "p95_latency_s": p95,
+ "per_query": per_query,
+ }
+
+
+def print_per_query_results(label_a, report_a, label_b, report_b):
+ """Print each golden query's expected titles and both backends' hit/miss + retrieved titles."""
+ print("\nPer-query results:")
+ for a, b in zip(report_a["per_query"], report_b["per_query"], strict=True):
+ expected = ", ".join(a["expected_titles"]) or "(unlabelled — excluded from recall)"
+ print(f" [{a['category'] or 'uncategorised'}] {a['query']}")
+ print(f" expected: {expected}")
+ for label, item in ((label_a, a), (label_b, b)):
+ status = {True: "HIT ", False: "MISS", None: "-- "}[item["hit"]]
+ titles = ", ".join(t for t in item["retrieved_titles"] if t is not None)
+ print(f" {label:26} {status} {item['latency_s']:.3f}s {titles}")
+
+
+def build_configs():
+ """Every (label, backend factory, strategy) to evaluate.
+
+ Postgres configurations for a chunk size with no complete batch are skipped
+ with a notice, so the eval still runs when only some batches are indexed.
+ """
+ configs = [("Pinecone semantic", LegacyPineconeDocsiteSearch, "semantic")]
+
+ for chunk_size in CHUNK_SIZES:
+ batch_id = resolve_batch_id(chunk_size)
+ if batch_id is None:
+ print(f"Skipping Postgres configs for chunk size {chunk_size} — no complete batch")
+ continue
+ for strategy in STRATEGIES:
+ label = f"Postgres {strategy} ({chunk_size})"
+ configs.append((label, lambda b=batch_id: DocsiteSearch(batch_id=b), strategy))
+
+ return configs
+
+
+def _fmt(recall):
+ """Render a recall figure, or a dash when the category had no scored queries."""
+ return "-" if recall is None else f"{recall:.2f}"
+
+
+def main():
+ with open(GOLDEN_QUERIES_PATH) as f:
+ golden_queries = yaml.safe_load(f)["queries"]
+
+ reports = {}
+ for label, make_backend, strategy in build_configs():
+ reports[label] = run_eval(golden_queries, make_backend, strategy)
+
+ print(f"\n{'configuration':<28} {'conceptual':>11} {'keyword':>9} {'overall':>9} {'p50':>8} {'p95':>8}")
+ for label, report in reports.items():
+ by_cat = report["recall_by_category"]
+ conceptual = by_cat.get("conceptual", {}).get("recall")
+ keyword = by_cat.get("keyword", {}).get("recall")
+ print(f"{label:<28} "
+ f"{_fmt(conceptual):>11} {_fmt(keyword):>9} {_fmt(report['recall_at_k']):>9} "
+ f"{report['p50_latency_s']:>7.3f}s {report['p95_latency_s']:>7.3f}s")
+
+ scored = next(iter(reports.values()))
+ print(f"\nScored {scored['queries_scored']} queries, skipped {scored['queries_skipped']}")
+
+ label_a, label_b = AGREEMENT_PAIR
+ if label_a in reports and label_b in reports:
+ agreement = compute_agreement(reports[label_a], reports[label_b])
+ print(f"\nMigration fidelity ({label_a} vs {label_b}):")
+ print(f" mean doc-title Jaccard={agreement['mean_jaccard']:.3f} "
+ f"across {len(agreement['per_query'])} queries")
+ for q in agreement["per_query"]:
+ print(f" {q['jaccard']:.2f} overlap={q['overlap']}/{q['union']} {q['query']}")
+ print_per_query_results(label_a, reports[label_a], label_b, reports[label_b])
+
+
+if __name__ == "__main__":
+ main()
diff --git a/services/search_docsite/tests/unit/test_docsite_search.py b/services/search_docsite/tests/unit/test_docsite_search.py
index 37630f90..c283481f 100644
--- a/services/search_docsite/tests/unit/test_docsite_search.py
+++ b/services/search_docsite/tests/unit/test_docsite_search.py
@@ -1,140 +1,253 @@
-"""Unit tests for DocsiteSearch — the Pinecone + OpenAI seam used by job_chat.
+"""Unit tests for the Postgres-backed DocsiteSearch (semantic/keyword/hybrid).
-These pin the contracts most exposed by the dependency bump (langchain-pinecone
-0.2.2→0.2.13, langchain-openai →1.x, pinecone 5→7):
-
- - the langchain `similarity_search_with_score(query=, k=, filter=)` signature
- and its `[(Document, score), ...]` return shape, consumed by `_semantic_search`
- - the pinecone `describe_index_stats().get("namespaces")` shape, consumed by
- `_get_most_recent_namespace`
- - that the module's dependency symbols still import under the new versions
-
-Every external boundary is mocked, so no network/credentials are touched (the
-repo-root conftest also blocks real anthropic/openai client construction here).
+get_db_connection and register_vector_type are mocked throughout, no real
+Postgres connection is made. The OpenAI embeddings client is mocked via the
+`_embeddings` attribute, matching the DocsiteIndexer test pattern.
"""
+import json
+from decimal import Decimal
from unittest.mock import MagicMock, patch
+import psycopg2
import pytest
-
import search_docsite.search_docsite as m
+from pgvector import Vector
from util import ApolloError
-class FakeDoc:
- """Stand-in for a langchain Document (page_content + metadata)."""
+def make_conn():
+ conn = MagicMock()
+ cur = MagicMock()
+ conn.cursor.return_value.__enter__.return_value = cur
+ return conn, cur
- def __init__(self, text, metadata=None):
- self.page_content = text
- self.metadata = metadata or {}
+def make_search(**kwargs):
+ ds = m.DocsiteSearch(**kwargs)
+ ds._embeddings = MagicMock()
+ ds._embeddings.embed_query.return_value = [0.1, 0.2, 0.3]
+ return ds
-def make_search(default_top_k=5):
- """Construct DocsiteSearch offline: collection_name given (skips the
- namespace lookup) and PineconeVectorStore patched (no real client)."""
- with patch.object(m, "PineconeVectorStore", return_value=MagicMock()):
- return m.DocsiteSearch(
- collection_name="docsite-20240101",
- default_top_k=default_top_k,
- embeddings=MagicMock(),
- )
+def patched(conn):
+ return patch.object(m, "get_db_connection", return_value=conn), patch.object(m, "register_vector_type")
-# --- _build_filter (pure logic) ------------------------------------------------
-@pytest.mark.parametrize(
- "kwargs, expected",
- [
- ({"doc_title": "Adaptor X"}, {"doc_title": {"$eq": "Adaptor X"}}),
- ({"docs_type": "general_docs"}, {"docs_type": {"$eq": "general_docs"}}),
- ],
-)
-def test_build_filter_single_key(kwargs, expected):
- ds = make_search()
- assert ds._build_filter(**kwargs) == expected
+# --- strategy dispatch -----------------------------------------------------
+
+def test_search_dispatches_to_semantic_strategy():
+ conn, _ = make_conn()
+ ds = make_search(batch_id=1)
+ with patched(conn)[0], patched(conn)[1], patch.object(ds, "_semantic_search", return_value=["r"]) as mock_sem:
+ result = ds.search("query", strategy="semantic")
+ assert result == ["r"]
+ mock_sem.assert_called_once()
+
+
+def test_search_raises_on_unknown_strategy():
+ conn, _ = make_conn()
+ ds = make_search(batch_id=1)
+ with patched(conn)[0], patched(conn)[1]:
+ with pytest.raises(ApolloError) as exc:
+ ds.search("query", strategy="nonsense")
+ assert exc.value.code == 400
+
+
+@pytest.mark.parametrize("strategy", ["hybrid", "keyword"])
+def test_search_rejects_threshold_for_non_semantic_strategies(strategy):
+ """RRF/FTS scores are not comparable to a cosine cutoff. Silently ignoring a
+ threshold here is a landmine: 0.8 against a 0.033-max score would drop
+ every result with no error."""
+ conn, _ = make_conn()
+ ds = make_search(batch_id=1)
+ with patched(conn)[0], patched(conn)[1]:
+ with pytest.raises(ApolloError) as exc:
+ ds.search("query", threshold=0.8, strategy=strategy)
+ assert exc.value.code == 400
+
+
+@pytest.mark.parametrize("strategy", ["hybrid", "keyword"])
+def test_search_allows_none_threshold_for_non_semantic_strategies(strategy):
+ conn, _ = make_conn()
+ ds = make_search(batch_id=1)
+ with patched(conn)[0], patched(conn)[1], \
+ patch.object(ds, "_keyword_search", return_value=["r"]), \
+ patch.object(ds, "_hybrid_search", return_value=["r"]):
+ assert ds.search("query", threshold=None, strategy=strategy) == ["r"]
-def test_build_filter_both_combines_with_and():
+# --- _resolve_current_batch --------------------------------------------------
+
+def test_resolve_current_batch_returns_newest_complete_batch_id():
+ conn, cur = make_conn()
+ cur.fetchone.return_value = (9,)
ds = make_search()
- assert ds._build_filter(doc_title="X", docs_type="general_docs") == {
- "$and": [{"doc_title": {"$eq": "X"}}, {"docs_type": {"$eq": "general_docs"}}]
- }
+ assert ds._resolve_current_batch(conn) == 9
-def test_build_filter_none_returns_none():
+def test_resolve_current_batch_raises_when_none_complete():
+ conn, cur = make_conn()
+ cur.fetchone.return_value = None
ds = make_search()
- assert ds._build_filter() is None
+ with pytest.raises(ApolloError) as exc:
+ ds._resolve_current_batch(conn)
+ assert exc.value.code == 404
-# --- _semantic_search (langchain return-shape contract) ------------------------
+# --- _semantic_search: (top_k, threshold) fallback semantics, ported from Pinecone tests ---
-def test_semantic_search_applies_threshold_and_passes_signature():
+def test_semantic_search_applies_threshold_and_falls_back_to_k_50():
+ conn, cur = make_conn()
+ cur.fetchall.return_value = [("a", "Doc A", "general_docs", 0.9), ("b", "Doc B", "general_docs", 0.4)]
ds = make_search()
- ds.vectorstore.similarity_search_with_score.return_value = [
- (FakeDoc("a"), 0.9),
- (FakeDoc("b"), 0.6),
- (FakeDoc("c"), 0.4), # below threshold, dropped
- ]
- results = ds._semantic_search(query="q", threshold=0.5)
+ results = ds._semantic_search(conn, batch_id=1, query="q", top_k=None, threshold=0.5, doc_title=None, docs_type=None)
- assert [r.score for r in results] == [0.9, 0.6]
- assert [r.text for r in results] == ["a", "b"]
- # Pin the langchain-pinecone call signature; threshold-only => k falls back to 50.
- ds.vectorstore.similarity_search_with_score.assert_called_once_with(
- query="q", k=50, filter=None
- )
+ assert [r.text for r in results] == ["a"]
+ params = cur.execute.call_args[0][1]
+ assert params["max_k"] == 50
def test_semantic_search_truncates_to_top_k_when_no_threshold():
+ conn, cur = make_conn()
+ cur.fetchall.return_value = [("a", "A", "t", 0.9), ("b", "B", "t", 0.8), ("c", "C", "t", 0.7)]
ds = make_search()
- ds.vectorstore.similarity_search_with_score.return_value = [
- (FakeDoc(t), s) for t, s in [("a", 0.9), ("b", 0.8), ("c", 0.7), ("d", 0.6)]
- ]
- results = ds._semantic_search(query="q", top_k=2)
+ results = ds._semantic_search(conn, batch_id=1, query="q", top_k=2, threshold=None, doc_title=None, docs_type=None)
assert [r.text for r in results] == ["a", "b"]
- ds.vectorstore.similarity_search_with_score.assert_called_once_with(
- query="q", k=2, filter=None
- )
def test_semantic_search_defaults_to_default_top_k():
+ conn, cur = make_conn()
+ cur.fetchall.return_value = [(str(i), str(i), "t", 0.9) for i in range(7)]
ds = make_search(default_top_k=5)
- ds.vectorstore.similarity_search_with_score.return_value = [
- (FakeDoc(str(i)), 0.9) for i in range(7)
- ]
- # Neither top_k nor threshold given => default_top_k (5) applies.
- results = ds._semantic_search(query="q")
+ results = ds._semantic_search(conn, batch_id=1, query="q", top_k=None, threshold=None, doc_title=None, docs_type=None)
assert len(results) == 5
- ds.vectorstore.similarity_search_with_score.assert_called_once_with(
- query="q", k=5, filter=None
- )
-# --- _get_most_recent_namespace (pinecone describe_index_stats shape) ----------
+# --- _keyword_search ---------------------------------------------------------
+
+def test_keyword_search_uses_ts_rank_and_returns_results():
+ conn, cur = make_conn()
+ cur.fetchall.return_value = [("a", "Doc A", "general_docs", 0.5)]
+ ds = make_search()
+
+ results = ds._keyword_search(conn, batch_id=1, query="webhook", top_k=None, doc_title=None, docs_type="general_docs")
+
+ assert len(results) == 1
+ assert results[0].text == "a"
+ sql = cur.execute.call_args[0][0]
+ assert "ts_rank_cd" in sql
+ assert "plainto_tsquery" in sql
+
+
+# --- _hybrid_search ------------------------------------------------------------
+
+def test_hybrid_search_runs_rrf_query_and_returns_results():
+ conn, cur = make_conn()
+ cur.fetchall.return_value = [("a", "Doc A", "general_docs", 0.032)]
+ ds = make_search()
+
+ results = ds._hybrid_search(conn, batch_id=1, query="webhook", top_k=5, doc_title=None, docs_type="general_docs")
+
+ assert len(results) == 1
+ sql = cur.execute.call_args[0][0]
+ assert "FULL OUTER JOIN" in sql
+ params = cur.execute.call_args[0][1]
+ assert params["candidate_k"] == 50
+ assert params["max_k"] == 5
-def _patch_pinecone(namespaces):
- index = MagicMock()
- index.describe_index_stats.return_value = {"namespaces": {ns: {} for ns in namespaces}}
- client = MagicMock()
- client.Index.return_value = index
- return patch.object(m, "Pinecone", return_value=client)
+
+def test_hybrid_search_score_is_json_serializable_float():
+ """Postgres returns RRF as `numeric`, which psycopg2 hands back as Decimal.
+ Decimal is not JSON-serializable, and entry.py's json.dump sits outside its
+ try/except — so this would kill the process, not return a 500."""
+ conn, cur = make_conn()
+ cur.fetchall.return_value = [("a", "Doc A", "general_docs", Decimal("0.032"))]
+ ds = make_search()
+
+ results = ds._hybrid_search(conn, batch_id=1, query="webhook", top_k=5, doc_title=None, docs_type="general_docs")
+
+ assert isinstance(results[0].score, float)
+ json.dumps(results[0].to_json()) # must not raise
-def test_get_most_recent_namespace_picks_latest_valid():
+def test_hybrid_search_casts_rrf_to_float8_in_sql():
+ """Belt and braces: the SQL itself must not produce numeric in the first place."""
+ conn, cur = make_conn()
+ cur.fetchall.return_value = []
ds = make_search()
- namespaces = ["docsite-20231231", "docsite-20240101", "other", "docsite-bad"]
- with _patch_pinecone(namespaces):
- assert ds._get_most_recent_namespace() == "docsite-20240101"
+ ds._hybrid_search(conn, batch_id=1, query="webhook", top_k=5, doc_title=None, docs_type=None)
+
+ sql = cur.execute.call_args[0][0]
+ assert "float8" in sql
+
+
+# --- missing schema ----------------------------------------------------------
-def test_get_most_recent_namespace_raises_when_none_valid():
+def test_search_maps_missing_pgvector_extension_to_503():
+ """register_vector raises ProgrammingError('vector type not found in the
+ database') when CREATE EXTENSION has not run. Since migrations moved to the
+ indexer, that is the first thing a reader hits on an un-indexed database."""
+ conn, _ = make_conn()
+ ds = make_search(batch_id=1)
+
+ with patch.object(m, "get_db_connection", return_value=conn), \
+ patch.object(m, "register_vector_type",
+ side_effect=psycopg2.ProgrammingError("vector type not found in the database")):
+ with pytest.raises(ApolloError) as exc:
+ ds.search("query")
+
+ assert exc.value.code == 503
+ assert "embed_docsite" in exc.value.message
+ conn.close.assert_called_once()
+
+
+def test_search_maps_missing_docsite_tables_to_503():
+ """Extension present, tables absent — the second way a database can be
+ un-indexed."""
+ conn, cur = make_conn()
+ cur.execute.side_effect = psycopg2.errors.UndefinedTable(
+ 'relation "docsite_batches" does not exist',
+ )
ds = make_search()
- with _patch_pinecone(["other", "docsite-bad", "docsite-2024"]):
+
+ with patch.object(m, "get_db_connection", return_value=conn), \
+ patch.object(m, "register_vector_type"):
with pytest.raises(ApolloError) as exc:
- ds._get_most_recent_namespace()
- assert exc.value.code == 404
+ ds.search("query")
+
+ assert exc.value.code == 503
+ assert "embed_docsite" in exc.value.message
+
+
+# --- vector binding ----------------------------------------------------------
+
+def test_semantic_search_binds_the_embedding_as_a_vector():
+ """psycopg2 has no adapter for list, so a list is sent as numeric[] and
+ `vector <=> numeric[]` resolves to no operator. Only Vector and ndarray are
+ registered by pgvector."""
+ conn, cur = make_conn()
+ cur.fetchall.return_value = []
+ ds = make_search()
+
+ ds._semantic_search(conn, batch_id=1, query="q", top_k=5, threshold=None, doc_title=None, docs_type=None)
+
+ params = cur.execute.call_args[0][1]
+ assert params["query_embedding"] == Vector([0.1, 0.2, 0.3])
+
+
+def test_hybrid_search_binds_the_embedding_as_a_vector():
+ conn, cur = make_conn()
+ cur.fetchall.return_value = []
+ ds = make_search()
+
+ ds._hybrid_search(conn, batch_id=1, query="q", top_k=5, doc_title=None, docs_type=None)
+
+ params = cur.execute.call_args[0][1]
+ assert params["query_embedding"] == Vector([0.1, 0.2, 0.3])
diff --git a/services/search_docsite/tests/unit/test_pinecone_legacy_search.py b/services/search_docsite/tests/unit/test_pinecone_legacy_search.py
new file mode 100644
index 00000000..ed67f6e6
--- /dev/null
+++ b/services/search_docsite/tests/unit/test_pinecone_legacy_search.py
@@ -0,0 +1,161 @@
+"""Unit tests for DocsiteSearch — the Pinecone + OpenAI seam used by job_chat.
+
+These pin the contracts most exposed by the dependency bump (langchain-pinecone
+0.2.2→0.2.13, langchain-openai →1.x, pinecone 5→7):
+
+ - the langchain `similarity_search_with_score(query=, k=, filter=)` signature
+ and its `[(Document, score), ...]` return shape, consumed by `_semantic_search`
+ - the pinecone `describe_index_stats().get("namespaces")` shape, consumed by
+ `_get_most_recent_namespace`
+ - that the module's dependency symbols still import under the new versions
+
+Every external boundary is mocked, so no network/credentials are touched (the
+repo-root conftest also blocks real anthropic/openai client construction here).
+"""
+
+from unittest.mock import MagicMock, patch
+
+import pytest
+import search_docsite.pinecone_legacy_search as m
+from util import ApolloError
+
+
+class FakeDoc:
+ """Stand-in for a langchain Document (page_content + metadata)."""
+
+ def __init__(self, text, metadata=None):
+ self.page_content = text
+ self.metadata = metadata or {}
+
+
+def make_search(default_top_k=5):
+ """Construct DocsiteSearch offline: collection_name given (skips the
+ namespace lookup) and PineconeVectorStore patched (no real client)."""
+ with patch.object(m, "PineconeVectorStore", return_value=MagicMock()):
+ return m.LegacyPineconeDocsiteSearch(
+ collection_name="docsite-20240101",
+ default_top_k=default_top_k,
+ embeddings=MagicMock(),
+ )
+
+
+# --- _build_filter (pure logic) ------------------------------------------------
+
+@pytest.mark.parametrize(
+ "kwargs, expected",
+ [
+ ({"doc_title": "Adaptor X"}, {"doc_title": {"$eq": "Adaptor X"}}),
+ ({"docs_type": "general_docs"}, {"docs_type": {"$eq": "general_docs"}}),
+ ],
+)
+def test_build_filter_single_key(kwargs, expected):
+ ds = make_search()
+ assert ds._build_filter(**kwargs) == expected
+
+
+def test_build_filter_both_combines_with_and():
+ ds = make_search()
+ assert ds._build_filter(doc_title="X", docs_type="general_docs") == {
+ "$and": [{"doc_title": {"$eq": "X"}}, {"docs_type": {"$eq": "general_docs"}}]
+ }
+
+
+def test_build_filter_none_returns_none():
+ ds = make_search()
+ assert ds._build_filter() is None
+
+
+# --- _semantic_search (langchain return-shape contract) ------------------------
+
+def test_semantic_search_applies_threshold_and_passes_signature():
+ ds = make_search()
+ ds.vectorstore.similarity_search_with_score.return_value = [
+ (FakeDoc("a"), 0.9),
+ (FakeDoc("b"), 0.6),
+ (FakeDoc("c"), 0.4), # below threshold, dropped
+ ]
+
+ results = ds._semantic_search(query="q", threshold=0.5)
+
+ assert [r.score for r in results] == [0.9, 0.6]
+ assert [r.text for r in results] == ["a", "b"]
+ # Pin the langchain-pinecone call signature; threshold-only => k falls back to 50.
+ ds.vectorstore.similarity_search_with_score.assert_called_once_with(
+ query="q", k=50, filter=None
+ )
+
+
+def test_semantic_search_truncates_to_top_k_when_no_threshold():
+ ds = make_search()
+ ds.vectorstore.similarity_search_with_score.return_value = [
+ (FakeDoc(t), s) for t, s in [("a", 0.9), ("b", 0.8), ("c", 0.7), ("d", 0.6)]
+ ]
+
+ results = ds._semantic_search(query="q", top_k=2)
+
+ assert [r.text for r in results] == ["a", "b"]
+ ds.vectorstore.similarity_search_with_score.assert_called_once_with(
+ query="q", k=2, filter=None
+ )
+
+
+def test_semantic_search_defaults_to_default_top_k():
+ ds = make_search(default_top_k=5)
+ ds.vectorstore.similarity_search_with_score.return_value = [
+ (FakeDoc(str(i)), 0.9) for i in range(7)
+ ]
+
+ # Neither top_k nor threshold given => default_top_k (5) applies.
+ results = ds._semantic_search(query="q")
+
+ assert len(results) == 5
+ ds.vectorstore.similarity_search_with_score.assert_called_once_with(
+ query="q", k=5, filter=None
+ )
+
+
+# --- _get_most_recent_namespace (pinecone describe_index_stats shape) ----------
+
+def _patch_pinecone(namespaces):
+ index = MagicMock()
+ index.describe_index_stats.return_value = {"namespaces": {ns: {} for ns in namespaces}}
+ client = MagicMock()
+ client.Index.return_value = index
+ return patch.object(m, "Pinecone", return_value=client)
+
+
+def test_get_most_recent_namespace_picks_latest_valid():
+ ds = make_search()
+ namespaces = ["docsite-20231231", "docsite-20240101", "other", "docsite-bad"]
+ with _patch_pinecone(namespaces):
+ assert ds._get_most_recent_namespace() == "docsite-20240101"
+
+
+def test_get_most_recent_namespace_raises_when_none_valid():
+ ds = make_search()
+ with _patch_pinecone(["other", "docsite-bad", "docsite-2024"]):
+ with pytest.raises(ApolloError) as exc:
+ ds._get_most_recent_namespace()
+ assert exc.value.code == 404
+
+
+# --- strategy guard -------------------------------------------------------------
+
+@pytest.mark.parametrize("strategy", ["hybrid", "keyword", "nonsense"])
+def test_legacy_search_raises_on_unsupported_strategy(strategy):
+ """Previously fell off the end of the method and returned None implicitly,
+ which surfaces downstream as a confusing TypeError. Reachable now that the
+ backend is selectable per request."""
+ ds = make_search()
+ with pytest.raises(ApolloError) as exc:
+ ds.search("query", strategy=strategy)
+ assert exc.value.code == 400
+
+
+def test_legacy_search_still_dispatches_semantic():
+ ds = make_search()
+ ds.vectorstore.similarity_search_with_score.return_value = [(FakeDoc("a"), 0.9)]
+
+ results = ds.search("query", strategy="semantic")
+
+ assert [r.text for r in results] == ["a"]
diff --git a/services/search_docsite/tests/unit/test_run_eval.py b/services/search_docsite/tests/unit/test_run_eval.py
new file mode 100644
index 00000000..02f62428
--- /dev/null
+++ b/services/search_docsite/tests/unit/test_run_eval.py
@@ -0,0 +1,189 @@
+"""Unit tests for the offline golden-query eval's scoring logic.
+Backends are fully faked — no real search/DB/network involved."""
+
+from unittest.mock import MagicMock
+
+import pytest
+from search_docsite.tests.eval.run_eval import compute_agreement, compute_recall_at_k, run_eval
+
+
+def test_compute_recall_at_k_true_when_any_expected_title_present():
+ assert compute_recall_at_k(["Doc A", "Doc B"], ["Doc B", "Doc C"]) is True
+
+
+def test_compute_recall_at_k_false_when_no_overlap():
+ assert compute_recall_at_k(["Doc A"], ["Doc Z"]) is False
+
+
+def make_fake_backend(titles_by_query):
+ """Fake backend whose .search() returns SearchResults with the given doc_titles per query."""
+ from embeddings.embeddings import SearchResult
+
+ def fake_search(query, top_k=None, strategy=None, docs_type=None):
+ return [SearchResult(f"text-{t}", {"doc_title": t, "docs_type": "general_docs"}, 0.9) for t in titles_by_query.get(query, [])]
+
+ backend = MagicMock()
+ backend.search.side_effect = fake_search
+ backend_cls = MagicMock(return_value=backend)
+ return backend_cls
+
+
+def test_run_eval_scores_labeled_queries_and_skips_unlabeled():
+ golden_queries = [
+ {"query": "how do I configure a webhook", "expected_doc_titles": ["Webhooks"]},
+ {"query": "what is a run", "expected_doc_titles": []}, # unlabeled — skipped from recall aggregate
+ ]
+ backend_cls = make_fake_backend({"how do I configure a webhook": ["Webhooks", "Other Doc"], "what is a run": ["Runs"]})
+
+ report = run_eval(golden_queries, backend_cls, strategy="hybrid", top_k=5)
+
+ assert report["queries_scored"] == 1
+ assert report["queries_skipped"] == 1
+ assert report["recall_at_k"] == 1.0
+ assert len(report["per_query"]) == 2
+
+
+def test_run_eval_computes_recall_across_multiple_labeled_queries():
+ golden_queries = [
+ {"query": "q1", "expected_doc_titles": ["A"]},
+ {"query": "q2", "expected_doc_titles": ["Z"]}, # backend won't return Z -> miss
+ ]
+ backend_cls = make_fake_backend({"q1": ["A"], "q2": ["B"]})
+
+ report = run_eval(golden_queries, backend_cls, strategy="semantic", top_k=5)
+
+ assert report["queries_scored"] == 2
+ assert report["recall_at_k"] == 0.5
+
+
+def test_run_eval_reports_latency_percentiles():
+ golden_queries = [{"query": "q1", "expected_doc_titles": ["A"]}]
+ backend_cls = make_fake_backend({"q1": ["A"]})
+
+ report = run_eval(golden_queries, backend_cls, strategy="hybrid", top_k=5)
+
+ assert "p50_latency_s" in report
+ assert "p95_latency_s" in report
+ assert report["p50_latency_s"] >= 0
+
+
+def test_compute_agreement_reports_perfect_overlap():
+ report_a = {"per_query": [{"query": "q", "retrieved_titles": ["A", "B"]}]}
+ report_b = {"per_query": [{"query": "q", "retrieved_titles": ["B", "A"]}]}
+
+ agreement = compute_agreement(report_a, report_b)
+
+ assert agreement["per_query"][0]["overlap"] == 2
+ assert agreement["per_query"][0]["jaccard"] == 1.0
+ assert agreement["mean_jaccard"] == 1.0
+
+
+def test_compute_agreement_reports_partial_overlap():
+ report_a = {"per_query": [{"query": "q", "retrieved_titles": ["A", "B"]}]}
+ report_b = {"per_query": [{"query": "q", "retrieved_titles": ["B", "C"]}]}
+
+ agreement = compute_agreement(report_a, report_b)
+
+ assert agreement["per_query"][0]["overlap"] == 1
+ assert agreement["per_query"][0]["union"] == 3
+ assert agreement["per_query"][0]["jaccard"] == pytest.approx(1 / 3)
+
+
+def test_compute_agreement_handles_both_backends_returning_nothing():
+ report_a = {"per_query": [{"query": "q", "retrieved_titles": []}]}
+ report_b = {"per_query": [{"query": "q", "retrieved_titles": []}]}
+
+ agreement = compute_agreement(report_a, report_b)
+
+ assert agreement["per_query"][0]["jaccard"] == 0.0
+ assert agreement["mean_jaccard"] == 0.0
+
+
+def test_compute_agreement_means_across_queries():
+ report_a = {"per_query": [
+ {"query": "q1", "retrieved_titles": ["A"]},
+ {"query": "q2", "retrieved_titles": ["X"]},
+ ]}
+ report_b = {"per_query": [
+ {"query": "q1", "retrieved_titles": ["A"]},
+ {"query": "q2", "retrieved_titles": ["Y"]},
+ ]}
+
+ agreement = compute_agreement(report_a, report_b)
+
+ assert agreement["mean_jaccard"] == pytest.approx(0.5)
+
+
+def test_run_eval_reports_recall_per_category():
+ golden_queries = [
+ {"query": "c1", "category": "conceptual", "expected_doc_titles": ["A"]},
+ {"query": "c2", "category": "conceptual", "expected_doc_titles": ["B"]},
+ {"query": "k1", "category": "keyword", "expected_doc_titles": ["C"]},
+ {"query": "k2", "category": "keyword", "expected_doc_titles": ["MISSING"]},
+ ]
+ backend_cls = make_fake_backend({"c1": ["A"], "c2": ["B"], "k1": ["C"], "k2": ["Z"]})
+
+ report = run_eval(golden_queries, backend_cls, strategy="semantic", top_k=5)
+
+ assert report["recall_by_category"]["conceptual"] == {"recall": 1.0, "scored": 2}
+ assert report["recall_by_category"]["keyword"] == {"recall": 0.5, "scored": 2}
+ assert report["recall_at_k"] == 0.75
+
+
+def test_run_eval_tolerates_queries_without_a_category():
+ golden_queries = [
+ {"query": "q1", "expected_doc_titles": ["A"]},
+ {"query": "q2", "category": "keyword", "expected_doc_titles": ["B"]},
+ ]
+ backend_cls = make_fake_backend({"q1": ["A"], "q2": ["B"]})
+
+ report = run_eval(golden_queries, backend_cls, strategy="semantic", top_k=5)
+
+ assert report["recall_at_k"] == 1.0
+ assert report["recall_by_category"] == {"keyword": {"recall": 1.0, "scored": 1}}
+
+
+def test_run_eval_excludes_unlabelled_queries_from_category_recall():
+ golden_queries = [
+ {"query": "k1", "category": "keyword", "expected_doc_titles": ["A"]},
+ {"query": "k2", "category": "keyword", "expected_doc_titles": []},
+ ]
+ backend_cls = make_fake_backend({"k1": ["A"], "k2": ["B"]})
+
+ report = run_eval(golden_queries, backend_cls, strategy="semantic", top_k=5)
+
+ assert report["recall_by_category"]["keyword"] == {"recall": 1.0, "scored": 1}
+ assert report["queries_skipped"] == 1
+
+
+def test_resolve_batch_id_returns_newest_matching_batch():
+ from unittest.mock import patch
+
+ import search_docsite.tests.eval.run_eval as m
+
+ conn = MagicMock()
+ cur = MagicMock()
+ conn.cursor.return_value.__enter__.return_value = cur
+ cur.fetchone.return_value = (11,)
+
+ with patch.object(m, "get_db_connection", return_value=conn):
+ assert m.resolve_batch_id(2500) == 11
+
+ assert cur.execute.call_args[0][1] == (2500,)
+ conn.close.assert_called_once()
+
+
+def test_resolve_batch_id_returns_none_when_no_batch_matches():
+ from unittest.mock import patch
+
+ import search_docsite.tests.eval.run_eval as m
+
+ conn = MagicMock()
+ cur = MagicMock()
+ conn.cursor.return_value.__enter__.return_value = cur
+ cur.fetchone.return_value = None
+
+ with patch.object(m, "get_db_connection", return_value=conn):
+ assert m.resolve_batch_id(1800) is None
+
+ conn.close.assert_called_once()
diff --git a/services/search_docsite/tests/unit/test_search_docsite_main.py b/services/search_docsite/tests/unit/test_search_docsite_main.py
new file mode 100644
index 00000000..c6724b84
--- /dev/null
+++ b/services/search_docsite/tests/unit/test_search_docsite_main.py
@@ -0,0 +1,98 @@
+"""Unit tests for search_docsite.main's backend selection and resolve_backend.
+
+The `backend` payload field lets the same query be run against either backend
+on demand for manual comparison.
+"""
+
+from unittest.mock import patch
+
+import pytest
+import search_docsite.search_docsite as m
+from util import ApolloError
+
+
+def test_main_defaults_to_pinecone_backend(monkeypatch):
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
+
+ with patch.object(m, "LegacyPineconeDocsiteSearch") as mock_legacy_cls:
+ mock_legacy_cls.return_value.search.return_value = []
+ m.main({"query": "webhooks"})
+
+ mock_legacy_cls.assert_called_once_with()
+ mock_legacy_cls.return_value.search.assert_called_once_with(query="webhooks")
+
+
+def test_main_payload_backend_overrides_env(monkeypatch):
+ monkeypatch.setenv("DOCSITE_SEARCH_BACKEND", "pinecone")
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
+
+ with patch.object(m, "DocsiteSearch") as mock_pg_cls:
+ mock_pg_cls.return_value.search.return_value = []
+ m.main({"query": "webhooks", "backend": "postgres"})
+
+ mock_pg_cls.return_value.search.assert_called_once_with(query="webhooks")
+
+
+def test_main_env_selects_postgres_when_no_payload_override(monkeypatch):
+ monkeypatch.setenv("DOCSITE_SEARCH_BACKEND", "postgres")
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
+
+ with patch.object(m, "DocsiteSearch") as mock_pg_cls:
+ mock_pg_cls.return_value.search.return_value = []
+ m.main({"query": "webhooks"})
+
+ mock_pg_cls.return_value.search.assert_called_once_with(query="webhooks")
+
+
+def test_main_rejects_unknown_backend(monkeypatch):
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
+
+ with pytest.raises(ApolloError) as exc:
+ m.main({"query": "webhooks", "backend": "sqlite"})
+
+ assert exc.value.code == 400
+
+
+def test_main_routes_index_params_per_backend(monkeypatch):
+ """The two classes take different constructor params. batch_id is meaningless
+ to Pinecone; collection_name is meaningless to Postgres."""
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
+
+ with patch.object(m, "DocsiteSearch") as mock_pg_cls:
+ mock_pg_cls.return_value.search.return_value = []
+ m.main({"query": "q", "backend": "postgres", "batch_id": 3, "collection_name": "ignored"})
+
+ mock_pg_cls.assert_called_once_with(batch_id=3)
+
+
+def test_main_routes_collection_name_to_pinecone_only(monkeypatch):
+ monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
+
+ with patch.object(m, "LegacyPineconeDocsiteSearch") as mock_legacy_cls:
+ mock_legacy_cls.return_value.search.return_value = []
+ m.main({"query": "q", "backend": "pinecone", "collection_name": "docsite-202501010000", "batch_id": 9})
+
+ mock_legacy_cls.assert_called_once_with(collection_name="docsite-202501010000")
+
+
+def test_resolve_backend_defaults_to_pinecone(monkeypatch):
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+ assert m.resolve_backend() is m.LegacyPineconeDocsiteSearch
+
+
+def test_resolve_backend_reads_env(monkeypatch):
+ monkeypatch.setenv("DOCSITE_SEARCH_BACKEND", "postgres")
+ assert m.resolve_backend() is m.DocsiteSearch
+
+
+def test_resolve_backend_override_wins_over_env(monkeypatch):
+ monkeypatch.setenv("DOCSITE_SEARCH_BACKEND", "pinecone")
+ assert m.resolve_backend("postgres") is m.DocsiteSearch
+
+
+def test_resolve_backend_rejects_unknown_name(monkeypatch):
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+ with pytest.raises(ApolloError) as exc:
+ m.resolve_backend("sqlite")
+ assert exc.value.code == 400
diff --git a/services/tools/search_documentation/search_documentation.py b/services/tools/search_documentation/search_documentation.py
index 609c5aa8..018cc172 100644
--- a/services/tools/search_documentation/search_documentation.py
+++ b/services/tools/search_documentation/search_documentation.py
@@ -5,17 +5,16 @@
1. As a standalone service via entry.py: bun py tools/search_documentation
2. As a tool by supervisor via search_documentation_tool()
"""
-import os
import sys
-from pathlib import Path
-from typing import Dict, List, Optional
from dataclasses import dataclass
+from pathlib import Path
+from typing import Dict
# Import utilities from services directory
sys.path.append(str(Path(__file__).parent.parent.parent))
-from util import create_logger, ApolloError
-from search_docsite.search_docsite import DocsiteSearch
+from search_docsite.search_docsite import resolve_backend
+from util import ApolloError, create_logger
logger = create_logger(__name__)
@@ -46,8 +45,10 @@ def _search_implementation(query: str, num_results: int) -> Dict:
"""
logger.info(f"Searching documentation for: {query[:100]}...")
+ # Both backends use semantic search with the same cosine-similarity cutoff, so
+ # results are directly comparable.
# Initialize docsite search
- docsite_search = DocsiteSearch()
+ docsite_search = resolve_backend()()
# Search with threshold for quality results
search_results = docsite_search.search(
@@ -91,7 +92,7 @@ def main(data: Dict) -> Dict:
raise
except Exception as e:
logger.exception("Error in search_documentation service")
- raise ApolloError(500, f"Documentation search failed: {str(e)}")
+ raise ApolloError(500, f"Documentation search failed: {str(e)}") from e
def search_documentation_tool(tool_input: Dict) -> str:
@@ -140,4 +141,4 @@ def search_documentation_tool(tool_input: Dict) -> str:
except Exception as e:
logger.exception("Error in search_documentation tool")
- raise ApolloError(500, f"Documentation search failed: {str(e)}")
+ raise ApolloError(500, f"Documentation search failed: {str(e)}") from e
diff --git a/services/tools/search_documentation/tests/unit/conftest.py b/services/tools/search_documentation/tests/unit/conftest.py
new file mode 100644
index 00000000..32aac060
--- /dev/null
+++ b/services/tools/search_documentation/tests/unit/conftest.py
@@ -0,0 +1,14 @@
+"""Test config for search_documentation unit tests.
+
+Importing `search_documentation.search_documentation` pulls in
+`search_docsite.pinecone_legacy_search`, whose module-level `OpenAIEmbeddings()`
+default arg validates credentials at construction (openai 2.x / langchain-openai 1.x).
+A key must therefore exist at import time.
+
+Dummy placeholders only.
+"""
+
+import os
+
+os.environ.setdefault("OPENAI_API_KEY", "sk-test-dummy")
+os.environ.setdefault("PINECONE_API_KEY", "pc-test-dummy")
diff --git a/services/tools/search_documentation/tests/unit/test_search_documentation.py b/services/tools/search_documentation/tests/unit/test_search_documentation.py
new file mode 100644
index 00000000..569f8af9
--- /dev/null
+++ b/services/tools/search_documentation/tests/unit/test_search_documentation.py
@@ -0,0 +1,20 @@
+"""Unit tests for search_documentation's delegation to the resolved backend."""
+
+from unittest.mock import patch
+
+import tools.search_documentation.search_documentation as m
+
+
+def test_search_implementation_applies_quality_gate_to_resolved_backend(monkeypatch):
+ """The 0.7 threshold is a quality gate on live traffic — it was silently
+ dropped once already."""
+ monkeypatch.delenv("DOCSITE_SEARCH_BACKEND", raising=False)
+
+ with patch.object(m, "resolve_backend") as mock_resolve:
+ backend = mock_resolve.return_value.return_value
+ backend.search.return_value = []
+ m._search_implementation("how do I use webhooks", 5)
+
+ backend.search.assert_called_once_with(
+ query="how do I use webhooks", top_k=5, threshold=0.7, strategy="semantic"
+ )