Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 8 additions & 5 deletions Makefile
Original file line number Diff line number Diff line change
@@ -1,20 +1,23 @@
.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:
flake8 src
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
1 change: 1 addition & 0 deletions src/app/models/documents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
22 changes: 22 additions & 0 deletions src/app/search/services/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -57,6 +58,7 @@ def __init__(self, client):
"document_corpus",
"document_desc",
"document_sdg",
"document_external_id",
"slice_content",
"slice_sdg",
"document_scrape_date",
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down
23 changes: 23 additions & 0 deletions src/app/services/sql_db/queries.py
Original file line number Diff line number Diff line change
Expand Up @@ -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="",
Expand All @@ -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(
Expand Down
4 changes: 3 additions & 1 deletion src/app/shared/infra/abst_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions src/app/tests/api/api_v1/test_search.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = [
Expand Down Expand Up @@ -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 = [
Expand Down
6 changes: 4 additions & 2 deletions src/app/tests/services/test_llm_proxy.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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"}

Expand Down
Loading