diff --git a/.env.example b/.env.example index 69cc2f31..a6fb8fca 100644 --- a/.env.example +++ b/.env.example @@ -22,6 +22,12 @@ AZURE_API_VERSION=0000-00-00-preview LLM_MODEL_NAME=fake-llm-model LLM_TEMPERATURE=0.7 +##### LANGSMITH TRACING #### +LANGSMITH_TRACING_ENABLED=False +LANGSMITH_PROJECT=fake_project +LANGSMITH_ENDPOINT=https://api.smith.langchain.com +LANGSMITH_API_KEY=FAKE_LANGSMITH_API_KEY + ##### POSTGRES ######### # PG_HOST=fake-prod-postgres.example.com PG_HOST=fake-dev-postgres.example.com diff --git a/poetry.lock b/poetry.lock index 02ca0669..31aabfb4 100644 --- a/poetry.lock +++ b/poetry.lock @@ -3033,33 +3033,40 @@ orjson = ">=3.11.5" [[package]] name = "langsmith" -version = "0.6.7" +version = "0.9.3" description = "Client library to connect to the LangSmith Observability and Evaluation Platform." optional = false python-versions = ">=3.10" groups = ["main", "metrics"] files = [ - {file = "langsmith-0.6.7-py3-none-any.whl", hash = "sha256:4bd4372b8bf724b86314f64644562b5598407614e04e74b536c09490d153bd61"}, - {file = "langsmith-0.6.7.tar.gz", hash = "sha256:d89c604a18fc606b7835d8e7924f7cdbe130ca2207bdff8f989590e50d65b802"}, + {file = "langsmith-0.9.3-py3-none-any.whl", hash = "sha256:aac85686868a7e61d77858b880230e510583c61ea48ac0a197e6ceff64ffb22b"}, + {file = "langsmith-0.9.3.tar.gz", hash = "sha256:868a007b3ecda002914b21212a953c77fcb1d71a8a09a94562c3b57941d3df3a"}, ] [package.dependencies] +anyio = ">=3.5.0" +distro = ">=1.7.0" httpx = ">=0.23.0,<1" orjson = {version = ">=3.9.14", markers = "platform_python_implementation != \"PyPy\""} packaging = ">=23.2" pydantic = ">=2,<3" requests = ">=2.0.0" requests-toolbelt = ">=1.0.0" +sniffio = ">=1.1" +typing-extensions = ">=4.0.0" uuid-utils = ">=0.12.0,<1.0" +websockets = ">=15.0" xxhash = ">=3.0.0" zstandard = ">=0.23.0" [package.extras] claude-agent-sdk = ["claude-agent-sdk (>=0.1.0) ; python_version >= \"3.10\""] +google-adk = ["google-adk (>=1.0.0)", "wrapt (>=1.16.0)"] langsmith-pyo3 = ["langsmith-pyo3 (>=0.1.0rc2)"] openai-agents = ["openai-agents (>=0.0.3)"] otel = ["opentelemetry-api (>=1.30.0)", "opentelemetry-exporter-otlp-proto-http (>=1.30.0)", "opentelemetry-sdk (>=1.30.0)"] pytest = ["pytest (>=7.0.0)", "rich (>=13.9.4)", "vcrpy (>=7.0.0)"] +strands-agents = ["opentelemetry-api (>=1.30.0)", "opentelemetry-exporter-otlp-proto-http (>=1.30.0)", "opentelemetry-sdk (>=1.30.0)", "strands-agents (>=0.1.0)", "strands-agents-tools (>=0.2.0)"] vcr = ["vcrpy (>=7.0.0)"] [[package]] @@ -7461,7 +7468,7 @@ version = "16.0" description = "An implementation of the WebSocket Protocol (RFC 6455 & 7692)" optional = false python-versions = ">=3.10" -groups = ["main"] +groups = ["main", "metrics"] files = [ {file = "websockets-16.0-cp310-cp310-macosx_10_9_universal2.whl", hash = "sha256:04cdd5d2d1dacbad0a7bf36ccbcd3ccd5a30ee188f2560b7a62a30d14107b31a"}, {file = "websockets-16.0-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:8ff32bb86522a9e5e31439a58addbb0166f0204d64066fb955265c4e214160f0"}, @@ -8160,4 +8167,4 @@ cffi = ["cffi (>=1.17,<2.0) ; platform_python_implementation != \"PyPy\" and pyt [metadata] lock-version = "2.1" python-versions = ">=3.10,<3.13" -content-hash = "333a1f8be5f702464913572c00f2dbf6bff16b9d8cb6fff265b530d2551dfb6f" +content-hash = "c948ba85567f29a2c5e410c31ae9ad28173f30802d82d2d6938c49d6439ea522" diff --git a/pyproject.toml b/pyproject.toml index 8cc42241..957f584b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,7 @@ langchain-mistralai = "^1.1.2" langchain-azure-ai = "^1.2.3" langgraph = "^1.1.10" mistralai = "^2.4.3" +langsmith = "^0.9.3" [tool.poetry.group.dev.dependencies] pytest = "^7.4.0" diff --git a/src/app/api/api_v1/endpoints/chat.py b/src/app/api/api_v1/endpoints/chat.py index 3c6ea5a3..bf7730be 100644 --- a/src/app/api/api_v1/endpoints/chat.py +++ b/src/app/api/api_v1/endpoints/chat.py @@ -59,6 +59,26 @@ } +def _build_agent_trace_context( + *, + endpoint: str, + session_id: UUID | None, + thread_id: UUID, + body: models.AgentContext, +) -> models.TraceContext: + env = settings.ENV + return { + "endpoint": endpoint, + "feature": "chat_agent", + "environment": env, + "session_id": str(session_id) if session_id else None, + "thread_id": str(thread_id), + "query_length": len(body.query or ""), + "sdg_filter": body.sdg_filter, + "corpora": list(body.corpora) if body.corpora else None, + } + + def get_params(body: models.Context) -> models.ContextOut: body.sources = body.sources[:7] @@ -400,6 +420,12 @@ async def agent_stream_response( try: session_id = extract_session_cookie(request) thread_id = _resolve_thread_id(body.thread_id) + trace_context = _build_agent_trace_context( + endpoint="/api/v1/qna/chat/agent_stream", + session_id=session_id, + thread_id=thread_id, + body=body, + ) if body.query is None: raise EmptyQueryError() @@ -415,6 +441,7 @@ async def agent_stream_response( data_collection=data_collection, session_id=session_id, thread_id=thread_id, + trace_context=trace_context, ), media_type="text/event-stream", headers=SSE_HEADERS, @@ -457,6 +484,13 @@ async def agent_response( logger.info("No thread_id provided. Generating new thread_id.") thread_id = uuid.uuid4() + trace_context = _build_agent_trace_context( + endpoint="/api/v1/qna/chat/agent", + session_id=session_id, + thread_id=thread_id, + body=body, + ) + if body.query is None: raise EmptyQueryError() @@ -480,6 +514,7 @@ async def agent_response( sdg_filter=body.sdg_filter, sp=sp, background_tasks=background_tasks, + trace_context=trace_context, ) else: res = await chatfactory.agent_message( @@ -487,6 +522,7 @@ async def agent_response( corpora=body.corpora, sdg_filter=body.sdg_filter, sp=sp, + trace_context=trace_context, ) all_docs = [] diff --git a/src/app/api/api_v1/endpoints/chat_utils.py b/src/app/api/api_v1/endpoints/chat_utils.py index 85adbabd..6f9db1bb 100644 --- a/src/app/api/api_v1/endpoints/chat_utils.py +++ b/src/app/api/api_v1/endpoints/chat_utils.py @@ -83,6 +83,7 @@ async def _stream_agent_with_memory( sp: SearchService, background_tasks: BackgroundTasks, thread_id: UUID, + trace_context: models.TraceContext | None = None, ) -> AsyncGenerator[dict[str, Any], None]: async with await psycopg.AsyncConnection[DictRow].connect( db_uri, @@ -103,6 +104,7 @@ async def _stream_agent_with_memory( sp=sp, background_tasks=background_tasks, streamed_ans=True, + trace_context=trace_context, ) async for chunk in stream: @@ -153,6 +155,7 @@ async def _stream_agent_response( data_collection: Any, session_id: UUID | None, thread_id: UUID, + trace_context: models.TraceContext | None = None, ) -> AsyncGenerator[str, None]: final_content = "" docs = None @@ -166,6 +169,7 @@ async def _stream_agent_response( sp=sp, background_tasks=background_tasks, thread_id=thread_id, + trace_context=trace_context, ) async for chunk in stream: diff --git a/src/app/core/config.py b/src/app/core/config.py index af5d9ba0..59c28570 100644 --- a/src/app/core/config.py +++ b/src/app/core/config.py @@ -50,6 +50,13 @@ def get_api_version(self) -> dict: LLM_TEMPERATURE: float ENV: str + # LANGSMITH / LANGCHAIN TRACING + # TBD: migrate Optional/Union annotations to PEP 604 shorthand (X | Y, X | None) in a dedicated refactor. + LANGSMITH_TRACING_ENABLED: bool = False + LANGSMITH_PROJECT: Optional[str] = None + LANGSMITH_ENDPOINT: Optional[str] = None + LANGSMITH_API_KEY: Optional[str] = None + # PG PG_USER: Optional[str] = None PG_PASSWORD: Optional[str] = None diff --git a/src/app/core/langsmith.py b/src/app/core/langsmith.py new file mode 100644 index 00000000..13682137 --- /dev/null +++ b/src/app/core/langsmith.py @@ -0,0 +1,45 @@ +import os + +from src.app.core.config import Settings +from src.app.utils.logger import logger as utils_logger + +logger = utils_logger(__name__) + + +def configure_langsmith_tracing(settings: Settings) -> None: + """Configure LangSmith tracing for LangChain-based model calls.""" + + is_enabled = settings.LANGSMITH_TRACING_ENABLED + + os.environ["LANGSMITH_TRACING"] = "true" if is_enabled else "false" + + if not is_enabled: + logger.info("langsmith_tracing_enabled=false") + return + + if settings.LANGSMITH_PROJECT: + os.environ["LANGCHAIN_PROJECT"] = settings.LANGSMITH_PROJECT + else: + logger.warning( + "langsmith_tracing_enabled=true but LANGSMITH_PROJECT is missing" + ) + + if settings.LANGSMITH_ENDPOINT: + os.environ["LANGCHAIN_ENDPOINT"] = settings.LANGSMITH_ENDPOINT + else: + logger.warning( + "langsmith_tracing_enabled=true but LANGSMITH_ENDPOINT is missing" + ) + + if settings.LANGSMITH_API_KEY: + os.environ["LANGSMITH_API_KEY"] = settings.LANGSMITH_API_KEY + else: + logger.warning( + "langsmith_tracing_enabled=true but LANGSMITH_API_KEY is missing" + ) + + logger.info( + "langsmith_tracing_enabled=true project=%s endpoint=%s", + settings.LANGSMITH_PROJECT, + settings.LANGSMITH_ENDPOINT, + ) diff --git a/src/app/core/lifespan.py b/src/app/core/lifespan.py index 2354a782..bdb5d4b6 100644 --- a/src/app/core/lifespan.py +++ b/src/app/core/lifespan.py @@ -5,6 +5,7 @@ from fastapi import FastAPI from qdrant_client import AsyncQdrantClient +from src.app.core.langsmith import configure_langsmith_tracing from src.app.shared.infra.llm_proxy import LLMProxy from src.app.shared.utils.dependencies import get_settings from src.app.tutor.service.tutor import init_chat_model @@ -13,6 +14,7 @@ @asynccontextmanager async def lifespan(app: FastAPI): settings = get_settings() + configure_langsmith_tracing(settings) await init_chat_model(settings) app.state.qdrant = AsyncQdrantClient( url=settings.QDRANT_HOST, diff --git a/src/app/models/chat.py b/src/app/models/chat.py index 097cc18f..fd2e53f1 100644 --- a/src/app/models/chat.py +++ b/src/app/models/chat.py @@ -69,6 +69,17 @@ class AgentContext(SDGFilter): corpora: tuple[str, ...] | None = None +class TraceContext(TypedDict): + endpoint: str + feature: str + environment: str + session_id: str | None + thread_id: str + query_length: int + sdg_filter: list[int] | None + corpora: list[str] | None + + class AgentResponse(BaseModel): content: str | None = None status: str | None = None diff --git a/src/app/shared/infra/abst_chat.py b/src/app/shared/infra/abst_chat.py index bc10b03b..c356ea46 100644 --- a/src/app/shared/infra/abst_chat.py +++ b/src/app/shared/infra/abst_chat.py @@ -39,6 +39,7 @@ stringify_docs_content, ) from src.app.shared.domain.exceptions import LanguageNotSupportedError +from src.app.shared.infra.tracing import TraceComponent from src.app.shared.utils.dependencies import get_settings from src.app.utils.decorators import log_time_and_error from src.app.utils.logger import log_environmental_impacts @@ -85,6 +86,21 @@ def __init__( }, } + def _build_non_agent_trace_context( + self, + operation: str, + **extra: Any, + ) -> dict[str, Any]: + settings = get_settings() + trace_context: dict[str, Any] = { + "component": TraceComponent.CHAT_NON_AGENT.value, + "operation": operation, + "environment": settings.ENV, + "model": getattr(self.chat_client, "model", None), + } + trace_context.update(extra) + return trace_context + @log_time_and_error async def json_formatter_agent(self, unformatted_input, expected_output): output = await self.chat_client.completion( @@ -101,6 +117,10 @@ async def json_formatter_agent(self, unformatted_input, expected_output): response_format={ "type": "json_object", }, + trace_context=self._build_non_agent_trace_context( + "json_formatter_agent", + has_expected_output=bool(expected_output), + ), ) json = extract_json_from_response(output) @@ -148,6 +168,10 @@ async def _detect_lang_with_llm(self, query: str) -> Dict[str, str]: response_format={ "type": "json_object", }, + trace_context=self._build_non_agent_trace_context( + "detect_language_with_llm", + query_length=len(query), + ), ) if isinstance(detected_lang, str): @@ -192,6 +216,11 @@ async def _detect_past_message_ref( }, ], response_format={"type": "json_object"}, + trace_context=self._build_non_agent_trace_context( + "detect_past_message_ref", + query_length=len(query), + history_length=len(history), + ), ) try: @@ -427,6 +456,11 @@ async def get_new_questions( + query, }, ], + trace_context=self._build_non_agent_trace_context( + "get_new_questions", + query_length=len(query), + history_length=len(history), + ), ) assert isinstance(res, str) @@ -472,11 +506,27 @@ async def rephrase_message( ] if streamed_ans: - res = self.chat_client.completion_stream(messages) + res = await self.chat_client.completion_stream( + messages, + trace_context=self._build_non_agent_trace_context( + "rephrase_message_stream", + query_length=len(message), + history_length=len(history), + docs_count=len(docs), + subject=subject, + ), + ) return self.get_stream_chunks(res) res = await self.chat_client.completion( messages=messages, + trace_context=self._build_non_agent_trace_context( + "rephrase_message", + query_length=len(message), + history_length=len(history), + docs_count=len(docs), + subject=subject, + ), ) return res @@ -521,11 +571,27 @@ async def chat_message( }, ] if streamed_ans: - res = await self.chat_client.completion_stream(messages) + res = await self.chat_client.completion_stream( + messages, + trace_context=self._build_non_agent_trace_context( + "chat_message_stream", + query_length=len(query), + history_length=len(history), + docs_count=len(docs), + subject=subject, + ), + ) return self.get_stream_chunks(res) res = await self.chat_client.completion( messages=messages, + trace_context=self._build_non_agent_trace_context( + "chat_message", + query_length=len(query), + history_length=len(history), + docs_count=len(docs), + subject=subject, + ), ) return res @@ -562,6 +628,7 @@ async def agent_message( sp: SearchService | None = None, background_tasks: BackgroundTasks | None = None, streamed_ans: bool = False, + trace_context: Optional[dict[str, Any]] = None, ): """ Sends a chat message handled by an agent. @@ -582,7 +649,27 @@ async def agent_message( agent_executor = await self._create_agent(memory=memory) + settings = get_settings() + + metadata: dict[str, Any] = { + "component": TraceComponent.CHAT_AGENT.value, + "environment": settings.ENV, + "thread_id": str(thread_id) if thread_id else None, + "corpora": list(corpora) if corpora else None, + "sdg_filter": sdg_filter, + } + + if trace_context: + metadata.update(trace_context) + + tags = ["welearn", "chat", "agent"] + endpoint = metadata.get("endpoint") + if endpoint: + tags.append(f"endpoint:{endpoint}") + config = RunnableConfig( + tags=tags, + metadata=metadata, configurable={ "thread_id": thread_id, "corpora": corpora, @@ -628,7 +715,13 @@ async def run_llm_with_json_parsing( model_class, fallback_formatter: str | None = None, ): - raw = await self.chat_client.completion(messages=messages) + raw = await self.chat_client.completion( + messages=messages, + trace_context=self._build_non_agent_trace_context( + "run_llm_with_json_parsing", + has_fallback_formatter=fallback_formatter is not None, + ), + ) if not isinstance(raw, str): raise ValueError("LLM response must be string") diff --git a/src/app/shared/infra/llm_proxy.py b/src/app/shared/infra/llm_proxy.py index 3a4be88c..32a3440f 100644 --- a/src/app/shared/infra/llm_proxy.py +++ b/src/app/shared/infra/llm_proxy.py @@ -1,12 +1,14 @@ from abc import ABC -from typing import Optional, Type, Union +from typing import Any, Optional, Type, Union import litellm from azure.ai.inference.aio import ChatCompletionsClient from azure.core.credentials import AzureKeyCredential +from langsmith import traceable from mistralai.client import Mistral from pydantic import BaseModel +from src.app.shared.infra.tracing import TRACE_RUN_TYPE_LLM, TraceName from src.app.utils.decorators import log_time_and_error from src.app.utils.logger import logger as utils_logger @@ -63,36 +65,75 @@ async def close_client(self): await self.client.close() @log_time_and_error + @traceable( + run_type=TRACE_RUN_TYPE_LLM, + name=TraceName.COMPLETION_NON_AGENT.value, + ) async def completion( self, messages: list, response_format: Optional[Union[dict, Type[BaseModel]]] = None, + trace_context: Optional[dict[str, Any]] = None, ) -> dict | str: - logger.info("starting completion with model_name=%s", self.model) + logger.info( + "starting completion with model_name=%s trace_context=%s", + self.model, + trace_context, + ) if self.is_azure_model: - return await self.az_completion(messages) + return await self.az_completion( + messages, + response_format=response_format, + trace_context=trace_context, + ) else: # We assume that if it's not an Azure model, it's a Mistral model for now. This can be extended in the future to support other types of models. - return await self.mistral_completion(messages) + return await self.mistral_completion( + messages, + response_format=response_format, + trace_context=trace_context, + ) - async def az_completion(self, messages: list): + @traceable( + run_type=TRACE_RUN_TYPE_LLM, + name=TraceName.AZURE_COMPLETION_NON_AGENT.value, + ) + async def az_completion( + self, + messages: list, + response_format: Optional[Union[dict, Type[BaseModel]]] = None, + trace_context: Optional[dict[str, Any]] = None, + ): if self.client is None: raise ValueError("Azure client is not initialized.") + completion_kwargs = {} + if response_format is not None: + completion_kwargs["response_format"] = response_format + response = await self.client.complete( messages=messages, max_tokens=2048, temperature=0.8, top_p=0.1, model=self.model, + **completion_kwargs, ) return response.choices[0].message.content - async def az_completion_stream(self, messages: list): + @traceable( + run_type=TRACE_RUN_TYPE_LLM, + name=TraceName.AZURE_COMPLETION_STREAM_NON_AGENT.value, + ) + async def az_completion_stream( + self, + messages: list, + trace_context: Optional[dict[str, Any]] = None, + ): if self.client is None: raise ValueError("Azure client is not initialized.") @@ -102,16 +143,77 @@ async def az_completion_stream(self, messages: list): return response - async def mistral_completion(self, messages: list): + @traceable( + run_type=TRACE_RUN_TYPE_LLM, + name=TraceName.COMPLETION_STREAM_NON_AGENT.value, + ) + async def completion_stream( + self, + messages: list, + trace_context: Optional[dict[str, Any]] = None, + ): + logger.info( + "starting completion_stream with model_name=%s trace_context=%s", + self.model, + trace_context, + ) + + if self.is_azure_model: + return await self.az_completion_stream( + messages, + trace_context=trace_context, + ) + + return await self.mistral_completion_stream( + messages, + trace_context=trace_context, + ) + + @traceable( + run_type=TRACE_RUN_TYPE_LLM, + name=TraceName.MISTRAL_COMPLETION_NON_AGENT.value, + ) + async def mistral_completion( + self, + messages: list, + response_format: Optional[Union[dict, Type[BaseModel]]] = None, + trace_context: Optional[dict[str, Any]] = None, + ): if self.client is None: raise ValueError("Mistral client is not initialized.") + completion_kwargs = {} + if response_format is not None: + completion_kwargs["response_format"] = response_format + response = await self.client.chat.complete_async( messages=messages, max_tokens=2048, temperature=0.8, top_p=0.1, model=self.model, + **completion_kwargs, ) return response.choices[0].message.content + + @traceable( + run_type=TRACE_RUN_TYPE_LLM, + name=TraceName.MISTRAL_COMPLETION_STREAM_NON_AGENT.value, + ) + async def mistral_completion_stream( + self, + messages: list, + trace_context: Optional[dict[str, Any]] = None, + ): + if self.client is None: + raise ValueError("Mistral client is not initialized.") + + response = await self.client.chat.stream_async( + messages=messages, + max_tokens=2048, + temperature=0.8, + top_p=0.1, + model=self.model, + ) + return response diff --git a/src/app/shared/infra/tracing.py b/src/app/shared/infra/tracing.py new file mode 100644 index 00000000..252266ea --- /dev/null +++ b/src/app/shared/infra/tracing.py @@ -0,0 +1,18 @@ +from enum import Enum + + +TRACE_RUN_TYPE_LLM = "llm" + + +class TraceComponent(str, Enum): + CHAT_NON_AGENT = "chat_non_agent" + CHAT_AGENT = "chat_agent" + + +class TraceName(str, Enum): + COMPLETION_NON_AGENT = "Completion (non-agent)" + AZURE_COMPLETION_NON_AGENT = "Azure completion (non-agent)" + AZURE_COMPLETION_STREAM_NON_AGENT = "Azure completion stream (non-agent)" + COMPLETION_STREAM_NON_AGENT = "Completion stream (non-agent)" + MISTRAL_COMPLETION_NON_AGENT = "Mistral completion (non-agent)" + MISTRAL_COMPLETION_STREAM_NON_AGENT = "Mistral completion stream (non-agent)" diff --git a/src/app/tests/services/test_llm_proxy.py b/src/app/tests/services/test_llm_proxy.py index d8492752..76bec9da 100644 --- a/src/app/tests/services/test_llm_proxy.py +++ b/src/app/tests/services/test_llm_proxy.py @@ -31,3 +31,87 @@ async def test_response_as_json_string(self): ) self.assertIsInstance(response, str) self.assertEqual(response, '{"key": "value"}') + + async def test_completion_forwards_response_format_to_mistral(self): + response_format = {"type": "json_object"} + with mock.patch.object( + self.proxy, "mistral_completion", new=AsyncMock(return_value="text") + ) as mistral_completion: + await self.proxy.completion( + messages=[{"role": "user", "content": "Hello"}], + response_format=response_format, + ) + + mistral_completion.assert_awaited_once_with( + [{"role": "user", "content": "Hello"}], + response_format=response_format, + trace_context=None, + ) + + async def test_completion_forwards_response_format_to_azure(self): + self.proxy.is_azure_model = True + response_format = {"type": "json_object"} + with mock.patch.object( + self.proxy, "az_completion", new=AsyncMock(return_value="text") + ) as az_completion: + await self.proxy.completion( + messages=[{"role": "user", "content": "Hello"}], + response_format=response_format, + ) + + az_completion.assert_awaited_once_with( + [{"role": "user", "content": "Hello"}], + response_format=response_format, + trace_context=None, + ) + + async def test_completion_stream_routes_to_mistral_and_forwards_trace_context(self): + messages = [{"role": "user", "content": "Hello"}] + trace_context = {"trace_id": "abc123"} + + with mock.patch.object( + self.proxy, + "mistral_completion_stream", + new=AsyncMock(return_value="mistral_stream"), + ) as mistral_completion_stream, mock.patch.object( + self.proxy, + "az_completion_stream", + new=AsyncMock(return_value="azure_stream"), + ) as az_completion_stream: + response = await self.proxy.completion_stream( + messages=messages, + trace_context=trace_context, + ) + + self.assertEqual(response, "mistral_stream") + mistral_completion_stream.assert_awaited_once_with( + messages, + trace_context=trace_context, + ) + az_completion_stream.assert_not_awaited() + + async def test_completion_stream_routes_to_azure_and_forwards_trace_context(self): + self.proxy.is_azure_model = True + messages = [{"role": "user", "content": "Hello"}] + trace_context = {"trace_id": "xyz789"} + + with mock.patch.object( + self.proxy, + "az_completion_stream", + new=AsyncMock(return_value="azure_stream"), + ) as az_completion_stream, mock.patch.object( + self.proxy, + "mistral_completion_stream", + new=AsyncMock(return_value="mistral_stream"), + ) as mistral_completion_stream: + response = await self.proxy.completion_stream( + messages=messages, + trace_context=trace_context, + ) + + self.assertEqual(response, "azure_stream") + az_completion_stream.assert_awaited_once_with( + messages, + trace_context=trace_context, + ) + mistral_completion_stream.assert_not_awaited() diff --git a/src/app/tutor/api/router.py b/src/app/tutor/api/router.py index 6d9dce82..76f5dd9d 100644 --- a/src/app/tutor/api/router.py +++ b/src/app/tutor/api/router.py @@ -162,7 +162,14 @@ async def create_syllabus( settings: Settings = Depends(get_settings), ) -> SyllabusResponse: session_id = extract_session_cookie(request) - results = await tutor_manager(body, lang, settings) + trace_context = { + "endpoint": request.url.path, + "feature": "syllabus_creation", + "session_id": str(session_id) if session_id else None, + "query_extracts_count": len(body.extracts), + "documents_count": len(body.documents), + } + results = await tutor_manager(body, lang, settings, trace_context=trace_context) # TODO: handle errors @@ -253,7 +260,19 @@ async def handle_syllabus_feedback( ] try: - syllabus = await chatfactory.chat_client.completion(messages=messages) + syllabus = await chatfactory.chat_client.completion( + messages=messages, + trace_context={ + "component": "tutor_syllabus_feedback", + "operation": "syllabus_feedback", + "environment": get_settings().ENV, + "endpoint": request.url.path, + "session_id": str(session_id) if session_id else None, + "feedback_length": len(body.feedback), + "documents_count": len(body.documents), + "extracts_count": len(body.extracts), + }, + ) if not isinstance(syllabus, str): raise ValueError("Syllabus feedback response is not a string") diff --git a/src/app/tutor/service/agents.py b/src/app/tutor/service/agents.py index 9d093ea6..33ee2580 100644 --- a/src/app/tutor/service/agents.py +++ b/src/app/tutor/service/agents.py @@ -1,10 +1,12 @@ import json import time from pathlib import Path +from typing import Any from langchain_core.language_models import BaseChatModel from langchain_core.output_parsers import StrOutputParser from langchain_core.prompts import ChatPromptTemplate +from langchain_core.runnables import RunnableConfig from src.app.shared.utils.utils import build_system_message from src.app.tutor.service.models import MessageWithResources, SyllabusResponseAgent @@ -38,8 +40,23 @@ def get_disciplinary_skills(): class TutorChatAgent: """Thin wrapper around a LangChain chat model with a fixed system prompt.""" - def __init__(self, name: str, model: BaseChatModel, system_prompt: str) -> None: - self.name = name + agent_name: str | None = None + agent_tag: str | None = None + + def __init__( + self, + model: BaseChatModel, + system_prompt: str, + trace_tags: list[str] | None = None, + trace_metadata: dict[str, Any] | None = None, + ) -> None: + self.name = self.agent_name or self.__class__.__name__ + self.trace_tags = list(trace_tags or []) + if self.agent_tag: + self.trace_tags.append(f"agent:{self.agent_tag}") + + self.trace_metadata = dict(trace_metadata or {}) + self.trace_metadata["agent"] = self.name prompt = ChatPromptTemplate.from_messages( [("system", system_prompt), ("human", "{user_prompt}")] ) @@ -48,7 +65,12 @@ def __init__(self, name: str, model: BaseChatModel, system_prompt: str) -> None: async def run(self, user_prompt: str) -> str: start_time = time.time() - response = await self.chain.ainvoke({"user_prompt": user_prompt}) + config = RunnableConfig( + tags=self.trace_tags, + metadata=self.trace_metadata, + run_name=f"Tutor ({self.name})", + ) + response = await self.chain.ainvoke({"user_prompt": user_prompt}, config=config) logger.debug( "agent_type=%s response_time=%s", self.name, time.time() - start_time ) @@ -58,7 +80,16 @@ async def run(self, user_prompt: str) -> str: class UniversityTeacherAgent(TutorChatAgent): """First-pass syllabus creation based on user documents and metadata.""" - def __init__(self, model: BaseChatModel, lang) -> None: + agent_name = "UniversityTeacherAgent" + agent_tag = "university_teacher" + + def __init__( + self, + model: BaseChatModel, + lang, + trace_tags: list[str] | None = None, + trace_metadata: dict[str, Any] | None = None, + ) -> None: system_prompt = build_system_message( role="the University Professor Agent, responsible for drafting the initial syllabus based on the course materials provided by the user. Your role is to structure the course content, ensuring it aligns with academic standards and effectively conveys the subject matter.", backstory="You are a highly experienced university professor with expertise in structuring academic courses. You understand the nuances of designing a syllabus that is comprehensive yet adaptable, providing a strong foundation for course delivery. Your experience spans multiple disciplines, and you excel at organizing complex information into a structured curriculum.", @@ -76,7 +107,12 @@ def __init__(self, model: BaseChatModel, lang) -> None: ), expected_output=f"You must follow this template :\n {TEMPLATES['template0']} and translate it into the target language: {lang}.", ) - super().__init__("UniversityTeacherAgent", model, system_prompt) + super().__init__( + model, + system_prompt, + trace_tags=trace_tags, + trace_metadata=trace_metadata, + ) async def generate(self, message: MessageWithResources) -> SyllabusResponseAgent: DISCIPLINARY_SKILLS = get_disciplinary_skills() @@ -97,8 +133,16 @@ async def generate(self, message: MessageWithResources) -> SyllabusResponseAgent class SDGExpertAgent(TutorChatAgent): """Injects sustainability and SDG alignment using WeLearn resources.""" + agent_name = "SDGExpertAgent" + agent_tag = "sdg_expert" + def __init__( - self, model: BaseChatModel, greencomp_competencies: str, lang: str + self, + model: BaseChatModel, + greencomp_competencies: str, + lang: str, + trace_tags: list[str] | None = None, + trace_metadata: dict[str, Any] | None = None, ) -> None: system_prompt = build_system_message( role="the Sustainability Expert Agent, responsible for integrating sustainability concepts into the syllabus. Your role is to ensure that the syllabus aligns with relevant sustainability principles and frameworks in a way that is appropriate for the course discipline.", @@ -117,7 +161,12 @@ def __init__( f"1. The revised syllabus, ready for the pedagogical engineer's review. You must follow this template :\n {TEMPLATES['template0']} and translate it into the target language: {lang}.." ), ) - super().__init__("SDGExpertAgent", model, system_prompt) + super().__init__( + model, + system_prompt, + trace_tags=trace_tags, + trace_metadata=trace_metadata, + ) self.greencomp_competencies = greencomp_competencies async def enhance( @@ -142,7 +191,17 @@ async def enhance( class PedagogicalEngineerAgent(TutorChatAgent): """Final polish focusing on pedagogy and GreenComp alignment.""" - def __init__(self, model: BaseChatModel, greencomp_competencies: str, lang) -> None: + agent_name = "PedagogicalEngineerAgent" + agent_tag = "pedagogical_engineer" + + def __init__( + self, + model: BaseChatModel, + greencomp_competencies: str, + lang, + trace_tags: list[str] | None = None, + trace_metadata: dict[str, Any] | None = None, + ) -> None: system_prompt = build_system_message( role="the Pedagogical Engineer Agent, responsible for ensuring that the syllabus adheres to best practices in pedagogy. Your role is to refine learning objectives, align assessments with learning outcomes, and include competencies from the EU GreenComp Framework. You optimize the syllabus for student engagement and effectiveness.", backstory="You are an experienced pedagogical engineer specializing in higher education course design. You are deeply familiar with competency-based learning and the EU GreenComp Framework, active learning strategies, and assessment alignment. Your expertise ensures that syllabi are not only well-structured but also effective for learning.", @@ -161,7 +220,12 @@ def __init__(self, model: BaseChatModel, greencomp_competencies: str, lang) -> N f"1. Final Syllabus: The polished syllabus, ready for user review. You must follow this template :\n {TEMPLATES['template0']} and translate it into the target language: {lang}." ), ) - super().__init__("PedagogicalEngineerAgent", model, system_prompt) + super().__init__( + model, + system_prompt, + trace_tags=trace_tags, + trace_metadata=trace_metadata, + ) self.greencomp_competencies = greencomp_competencies async def refine(self, syllabus: SyllabusResponseAgent) -> SyllabusResponseAgent: diff --git a/src/app/tutor/service/tutor.py b/src/app/tutor/service/tutor.py index 5d7f2203..0853c742 100644 --- a/src/app/tutor/service/tutor.py +++ b/src/app/tutor/service/tutor.py @@ -63,7 +63,10 @@ async def close_chat_model() -> None: async def tutor_manager( - content: TutorSyllabusRequest, lang: str, settings: Settings + content: TutorSyllabusRequest, + lang: str, + settings: Settings, + trace_context: dict | None = None, ) -> list[SyllabusResponseAgent]: formatted_content = MessageWithResources( lang=lang, @@ -83,10 +86,39 @@ async def tutor_manager( "Chat model not initialized. Call init_chat_model() at startup." ) - teacher_agent = UniversityTeacherAgent(chat_model, lang) - sdg_agent = SDGExpertAgent(chat_model, GREENCOMP_COMPETENCIES, lang) + base_tags = ["welearn", "tutor", "syllabus"] + base_metadata = { + "component": "tutor_syllabus", + "environment": settings.ENV, + "language": lang, + "course_title": content.course_title, + } + + if trace_context: + base_metadata.update(trace_context) + endpoint = trace_context.get("endpoint") + if endpoint: + base_tags.append(f"endpoint:{endpoint}") + + teacher_agent = UniversityTeacherAgent( + chat_model, + lang, + trace_tags=base_tags, + trace_metadata=base_metadata, + ) + sdg_agent = SDGExpertAgent( + chat_model, + GREENCOMP_COMPETENCIES, + lang, + trace_tags=base_tags, + trace_metadata=base_metadata, + ) pedagogical_agent = PedagogicalEngineerAgent( - chat_model, GREENCOMP_COMPETENCIES, lang + chat_model, + GREENCOMP_COMPETENCIES, + lang, + trace_tags=base_tags, + trace_metadata=base_metadata, ) teacher_response = await teacher_agent.generate(formatted_content)