diff --git a/Makefile b/Makefile index 8f595567..c5aa088d 100644 --- a/Makefile +++ b/Makefile @@ -1,11 +1,14 @@ -.PHONY: run-dev +.PHONY: run-dev generate-baml + +generate-baml: + baml-cli generate --from ./src/app/baml_src run-poetry: - poetry run baml-cli generate --from ./src/app/baml_src + generate-baml poetry run uvicorn src.main:app --reload run-dev: - baml-cli generate --from ./src/app/baml_src + generate-baml uvicorn src.main:app --reload lint: @@ -13,8 +16,8 @@ lint: isort src black src -test-poetry: +test-poetry: generate-baml poetry run pytest -s -v --cov=src --cov-report=term-missing --cov-fail-under=82 --cov-report=html -test: +test: generate-baml pytest -s -v --cov=src --cov-report=term-missing --cov-fail-under=82 --cov-report=html diff --git a/src/app/models/documents.py b/src/app/models/documents.py index c9840c91..70473d26 100644 --- a/src/app/models/documents.py +++ b/src/app/models/documents.py @@ -22,6 +22,7 @@ class DocumentPayloadModel(BaseModel): document_sdg: list[int] document_title: str document_url: str + document_external_id: str | None = None slice_content: str slice_sdg: int | None diff --git a/src/app/search/services/search.py b/src/app/search/services/search.py index ae485268..8edccafb 100644 --- a/src/app/search/services/search.py +++ b/src/app/search/services/search.py @@ -26,6 +26,7 @@ from src.app.services.helpers import convert_embedding_bytes from src.app.services.sql_db.queries import ( get_embeddings_model_id_according_name, + get_external_ids_by_document_ids_sync, get_subject, ) from src.app.shared.domain.exceptions import CollectionNotFoundError, ModelNotFoundError @@ -57,6 +58,7 @@ def __init__(self, client): "document_corpus", "document_desc", "document_sdg", + "document_external_id", "slice_content", "slice_sdg", "document_scrape_date", @@ -315,6 +317,8 @@ async def search_handler( ex, ) + sorted_data = await self._enrich_with_external_ids(sorted_data) + if without_vectors: points_without_vectors = [ point.model_copy(update={"vector": None}) for point in sorted_data @@ -323,6 +327,24 @@ async def search_handler( return sorted_data + async def _enrich_with_external_ids( + self, points: list[http_models.ScoredPoint] + ) -> list[http_models.ScoredPoint]: + document_ids = [ + str(point.payload["document_id"]) + for point in points + if point.payload and point.payload.get("document_id") + ] + external_ids = await run_in_threadpool( + get_external_ids_by_document_ids_sync, document_ids + ) + for point in points: + if point.payload: + point.payload["document_external_id"] = external_ids.get( + str(point.payload.get("document_id")) + ) + return points + @log_time_and_error async def search_group_by_document( self, diff --git a/src/app/services/sql_db/queries.py b/src/app/services/sql_db/queries.py index 55ef1803..df851f0a 100644 --- a/src/app/services/sql_db/queries.py +++ b/src/app/services/sql_db/queries.py @@ -151,6 +151,7 @@ def get_documents_payload_by_ids_sync(documents_ids: list[str]) -> list[Document document_title=doc.title, document_url=doc.url, document_desc=doc.description, + document_external_id=doc.external_id, document_sdg=[sdg[0] for sdg in short_sdg_list], document_details=doc.details, slice_content="", @@ -163,6 +164,28 @@ def get_documents_payload_by_ids_sync(documents_ids: list[str]) -> list[Document return docs +def get_external_ids_by_document_ids_sync(document_ids: list[str]) -> dict[str, str]: + """ + Get the external_id of documents by their ids. + Args: + document_ids: The list of document ids to look up. + + Returns: + A mapping of document_id (str) to external_id. + """ + if not document_ids: + return {} + + with session_maker() as s: + rows = s.execute( + select(WeLearnDocument.id, WeLearnDocument.external_id).where( + WeLearnDocument.id.in_(document_ids) + ) + ).all() + + return {str(row.id): row.external_id for row in rows if row.external_id} + + def register_endpoint(endpoint, session_id, http_code): with session_maker() as session: endpoint_request = EndpointRequest( diff --git a/src/app/shared/infra/abst_chat.py b/src/app/shared/infra/abst_chat.py index 7d22e56d..da15ab89 100644 --- a/src/app/shared/infra/abst_chat.py +++ b/src/app/shared/infra/abst_chat.py @@ -575,7 +575,9 @@ async def syllabus_feedback_completion( messages: list[dict], max_tokens: int, ) -> str: - result = await self.chat_client.completion(messages=messages, max_tokens=max_tokens) + result = await self.chat_client.completion( + messages=messages, max_tokens=max_tokens + ) if not isinstance(result, str): raise ValueError("Syllabus feedback response is not a string") return result diff --git a/src/app/tests/api/api_v1/test_search.py b/src/app/tests/api/api_v1/test_search.py index e1185d46..33aaa074 100644 --- a/src/app/tests/api/api_v1/test_search.py +++ b/src/app/tests/api/api_v1/test_search.py @@ -455,6 +455,7 @@ def test_documents_by_ids_single_doc(self, session_maker_mock, *mocks): id=doc_id, description="Desc", details={"k": "v"}, + external_id="external-id-123", ) session.query.return_value.where.return_value.options.return_value.all.return_value = [ @@ -515,6 +516,7 @@ def test_documents_by_ids_corpus_missing(self, session_maker_mock, *mocks): id=doc_id, description="Desc", details={}, + external_id=None, ) session.query.return_value.where.return_value.options.return_value.all.return_value = [ diff --git a/src/app/tests/services/test_llm_proxy.py b/src/app/tests/services/test_llm_proxy.py index 9b4090f4..e0590c29 100644 --- a/src/app/tests/services/test_llm_proxy.py +++ b/src/app/tests/services/test_llm_proxy.py @@ -1,7 +1,7 @@ import unittest +from types import SimpleNamespace from unittest import mock from unittest.mock import AsyncMock -from types import SimpleNamespace from src.app.shared.infra.llm_proxy import LLMProxy @@ -158,7 +158,9 @@ async def test_azure_completion_records_langsmith_usage_on_current_run(self): {"ls_provider": "azure", "ls_model_name": "fake_model"} ) - async def test_record_langsmith_usage_does_not_overwrite_existing_provider_metadata(self): + async def test_record_langsmith_usage_does_not_overwrite_existing_provider_metadata( + self, + ): run_tree = mock.Mock() run_tree.metadata = {"ls_provider": "existing", "ls_model_name": "preset"}