diff --git a/README.md b/README.md index fe765910..97c92123 100644 --- a/README.md +++ b/README.md @@ -42,6 +42,12 @@ Tools used for development: Start the Postgres, Django REST, and React services by starting Docker Desktop and running `docker compose up --build` +Local development servers: + +- **Frontend (React dev server):** http://localhost:3000 — when developing locally, open the React app at this address. +- **Backend / API (Django):** http://localhost:8000 + - *Note:* if you open http://localhost:8000 in a browser you will see a minimal single-page fallback for the frontend, which is non-functional for development. To view the working frontend during development, use http://localhost:3000. + #### Postgres The application supports connecting to PostgreSQL databases via: diff --git a/config/env/dev.env.example b/config/env/dev.env.example index b8e195cf..7960c654 100644 --- a/config/env/dev.env.example +++ b/config/env/dev.env.example @@ -10,13 +10,14 @@ SQL_PASSWORD=balancer # Connection Type Examples: # -# CloudNativePG (Kubernetes service within cluster): -# SQL_HOST=balancer-postgres-rw -# SQL_HOST=balancer-postgres-rw.balancer.svc.cluster.local +# CloudNativePG (primary — Philthy Civic Cloud shared cluster): +# SQL_HOST=shared-cluster-rw.cloudnative-pg.svc.cluster.local +# SQL_DATABASE=balancer # (SSL typically not required within cluster) # -# AWS RDS (External database): +# AWS RDS (legacy — for migration contexts): # SQL_HOST=balancer-db.xxxxx.us-east-1.rds.amazonaws.com +# SQL_DATABASE=balancer_dev # (SSL typically required - set SQL_SSL_MODE if needed) # # Local development: diff --git a/deploy/manifests/balancer/base/secret.template.yaml b/deploy/manifests/balancer/base/secret.template.yaml index e003a6ce..42b86d9a 100644 --- a/deploy/manifests/balancer/base/secret.template.yaml +++ b/deploy/manifests/balancer/base/secret.template.yaml @@ -6,6 +6,9 @@ # repository. Secrets should be created in each target cluster using cluster-specific # tools (e.g., SealedSecrets in the cfp-sandbox-cluster). # +# CloudNativePG is the primary database target (Philthy Civic Cloud shared cluster). +# AWS RDS was used previously and may still be referenced for migration contexts. +# apiVersion: v1 kind: Secret metadata: @@ -20,8 +23,10 @@ stringData: REACT_APP_API_BASE_URL: https://balancer.sandbox.k8s.phl.io/ SECRET_KEY: randomly_generated_key_ere SQL_ENGINE: django.db.backends.postgresql - SQL_HOST: sql_host_here + # CloudNativePG (primary): shared-cluster-rw.cloudnative-pg.svc.cluster.local + # AWS RDS (legacy): balancer-db.xxxxx.us-east-1.rds.amazonaws.com + SQL_HOST: shared-cluster-rw.cloudnative-pg.svc.cluster.local SQL_PORT: '5432' - SQL_DATABASE: balancer_dev + SQL_DATABASE: balancer SQL_USER: balancer SQL_PASSWORD: sql_password_here diff --git a/frontend/src/components/Header/MdNavBar.tsx b/frontend/src/components/Header/MdNavBar.tsx index 550b74d2..7a12222a 100644 --- a/frontend/src/components/Header/MdNavBar.tsx +++ b/frontend/src/components/Header/MdNavBar.tsx @@ -89,10 +89,10 @@ const MdNavBar = (props: LoginFormProps) => {
  • - Medical Suggester + Medication Suggester
  • diff --git a/frontend/src/pages/PatientManager/NewPatientForm.tsx b/frontend/src/pages/PatientManager/NewPatientForm.tsx index 94c718de..ebdf40ac 100644 --- a/frontend/src/pages/PatientManager/NewPatientForm.tsx +++ b/frontend/src/pages/PatientManager/NewPatientForm.tsx @@ -460,6 +460,7 @@ const NewPatientForm = ({ id="suicide" name="suicide" type="radio" + value="No" checked={newPatientInfo.Suicide === "No"} onChange={(e) => handleRadioChange(e, "Suicide")} className="w-4 h-4 text-indigo-600 border-gray-300 focus:ring-indigo-600" diff --git a/server/api/views/assistant/agentic_loop.py b/server/api/views/assistant/agentic_loop.py new file mode 100644 index 00000000..32b3f50c --- /dev/null +++ b/server/api/views/assistant/agentic_loop.py @@ -0,0 +1,153 @@ +import json +import logging + +from api.views.assistant.assistant_types import ( + AgentResult, + ToolCallExecution, + ToolCallStatus, + TokenUsage, +) + +logger = logging.getLogger(__name__) + + +def run_agentic_loop( + response, client, model_defaults: dict, tools: list, user +) -> AgentResult: + + # Every tool call the agentic loop made before exiting + agentic_loop_tool_call_executions= [] + # Token usage for every responses.create call, one entry per iteration + agentic_loop_token_usage: list[TokenUsage] = [] + + # TODO: Cap the number of iterations — a model that keeps calling tools never exits, hanging the request and adding OpenAI cost + while True: + # At the top of the body, so the initial response and the terminal iteration are each counted exactly once + # get_token_usage never raises: it runs on the web request path, so an unrecognized usage shape must not fail a user's request. + agentic_loop_token_usage.append(get_token_usage(response)) + + # TODO: Add a schema function to ToolCallExecution + # user is threaded through so tools that need it get it at dispatch time + tool_output_schemas, tool_call_executions = handle_tool_calls(response, tools, user) + + # TODO: Rewrite to .append every iteration's list of tools or decide whether + # to add iteration integer to ToolCallExecution to split the flat tool_calls into iterations + + # TODO: Add a data type to contain each iteration's parameters, response id, + # token usage, and tool calls or output text from the client response and a function + # for the output text and corresponding response id + + # .extend splices every iteration's list of tools into one list + agentic_loop_tool_call_executions.extend(tool_call_executions) + + # Exit agentic loop when model response doesn't contain any tool calls + if not tool_output_schemas: + return AgentResult( + output_text=response.output_text, + response_id=response.id, + tool_calls=agentic_loop_tool_call_executions, + token_usages=agentic_loop_token_usage, + ) + + # TODO: Add error handling to collect partial AgentResult tool calls and token usage — + # decide first whether AgentResult describes only successful runs or whatever happened + response = client.responses.create( + input=tool_output_schemas, + previous_response_id=response.id, + **model_defaults, + ) + +def get_token_usage(response) -> TokenUsage: + """Token usage for one response""" + + # Guard the whole chain when usage is missing because this also runs on the web request path + # Field names from openai ResponseUsage, checked against 2.29.0. requirements.txt doesn't pin openai, + # and nothing here checks types, so an SDK upgrade could blank or change these silently + + usage = getattr(response, "usage", None) + input_details = getattr(usage, "input_tokens_details", None) + output_details = getattr(usage, "output_tokens_details", None) + + return TokenUsage( + input_tokens=getattr(usage, "input_tokens", None), + cached_input_tokens=getattr(input_details, "cached_tokens", None), + output_tokens=getattr(usage, "output_tokens", None), + reasoning_output_tokens=getattr(output_details, "reasoning_tokens", None), + ) + + +def handle_tool_calls( + response, tools: list, user +) -> tuple[list[dict], list[ToolCallExecution]]: + + # Index the tools by name so a model-supplied call name can be looked up. .get() + # returns None for an unknown name, handled explicitly below. + tools_by_name = {tool.name: tool for tool in tools} + + tool_output_schemas = [] + tool_call_executions: list[ToolCallExecution] = [] + + for response_item in response.output: + if response_item.type == "reasoning": + #logger.info(f"Reasoning step: {response_item.summary}") + pass + + elif response_item.type == "function_call": + + tool_output, tool_call_execution = _execute_function_call(response_item, tools_by_name, user) + + tool_output_schemas.append( + { + "type": "function_call_output", + "call_id": response_item.call_id, + "output": tool_output, + } + ) + + tool_call_executions.append(tool_call_execution) + + + return tool_output_schemas, tool_call_executions + + +def _execute_function_call( + response_item, tools_by_name: dict, user +) -> tuple[str, ToolCallExecution]: + + target_tool = tools_by_name.get(response_item.name) + + # Parsed below; stays None if the model's argument JSON can't be parsed, + # so a FAILED record still reports whatever we managed to read. + arguments = None + + if target_tool is None: + msg = f"ERROR - No tool registered for function call: {response_item.name}" + logger.error(msg) + return msg, ToolCallExecution( + name=response_item.name, + status=ToolCallStatus.UNREGISTERED, + error=msg, + ) + + try: + arguments = json.loads(response_item.arguments) + logger.info( + f"Invoking tool: {response_item.name} with arguments: {arguments}" + ) + tool_output = target_tool.run(user=user, **arguments) + logger.info(f"Tool {response_item.name} completed successfully") + return tool_output, ToolCallExecution( + name=response_item.name, + status=ToolCallStatus.OK, + arguments=arguments, + output=tool_output, + ) + except Exception as e: + msg = f"Error executing function call: {response_item.name}: {e}" + logger.error(msg, exc_info=True) + return msg, ToolCallExecution( + name=response_item.name, + status=ToolCallStatus.FAILED, + arguments=arguments, + error=str(e), + ) diff --git a/server/api/views/assistant/assistant_prompts.py b/server/api/views/assistant/assistant_prompts.py index 44bf9b9b..7ce454b0 100644 --- a/server/api/views/assistant/assistant_prompts.py +++ b/server/api/views/assistant/assistant_prompts.py @@ -1,3 +1,15 @@ + +# TODO: Replace the {name}/{page_number} citation template with a filled-in example, +# e.g. [Name advancespharmaco.pdf, Page 9], and require exactly one page per citation. +# Note both known importers pass their string through verbatim (no .format() reads the braces) +# The only thing interpreting the braces is the model +# This and the UUID in search_tool.py block the eval's citation scoring +# Citations are unparseable until both land, which blocks the eval's scoring layer: citation +# accuracy is the cheapest real signal available, and a parser written before these two +# fixes would measure prompt drift rather than accuracy. + +# TODO: When ask_database is registered again, mention it here — the prompt names only +# search_documents and says to "ALWAYS use" it first, steering the model away from ask_database INSTRUCTIONS = """ You are an AI assistant that helps users find and understand information about bipolar disorder from your internal library of bipolar disorder research sources using semantic search. diff --git a/server/api/views/assistant/assistant_services.py b/server/api/views/assistant/assistant_services.py index ac339b9f..2c91708a 100644 --- a/server/api/views/assistant/assistant_services.py +++ b/server/api/views/assistant/assistant_services.py @@ -3,65 +3,35 @@ from openai import OpenAI -from .assistant_prompts import INSTRUCTIONS -from .tool_services import ( - SEARCH_TOOLS_SCHEMA, - make_search_tool_mapping, - handle_tool_calls_with_reasoning, -) +from api.views.assistant.assistant_prompts import INSTRUCTIONS +from api.views.assistant.tool_services import TOOLS +from api.views.assistant.assistant_types import AgentResult +from api.views.assistant.agentic_loop import run_agentic_loop logger = logging.getLogger(__name__) +# Module-level so eval_assistant.py can import it and log which model the run used +MODEL_NAME = "gpt-5-nano" + def run_assistant( - message: str, user, + message: str, previous_response_id: str | None = None, -) -> tuple[str, str]: - """Wire together the OpenAI client, retrieval, and the agentic reasoning loop. - - Parameters - ---------- - message : str - The user's input message. - user : User - The Django user object used for document access control in search_documents. - previous_response_id : str | None - ID of a prior response for multi-turn conversation continuity. - - Returns - ------- - tuple[str, str] - (final_response_output_text, final_response_id) - """ - # TODO: Track total duration, cost metrics, and tool_calls_made count - # and return them from run_assistant for use in eval_assistant.py CSV output +) -> AgentResult: client = OpenAI(api_key=os.environ.get("OPENAI_API_KEY")) MODEL_DEFAULTS = { "instructions": INSTRUCTIONS, - "model": "gpt-5-nano", # 400,000 token context window - # A summary of the reasoning performed by the model. This can be useful for debugging and understanding the model's reasoning process. + "model": MODEL_NAME, + # TODO: Flip "summary" to "auto" once this org is confirmed verified with OpenAI "reasoning": {"effort": "low", "summary": None}, - "tools": SEARCH_TOOLS_SCHEMA, + "tools": [tool.schema() for tool in TOOLS], } - # TOOLS_SCHEMA tells the model what tools exist and what arguments to generate. - # tool_mapping wires those tool names to the Python functions that execute them. - # They are separate because the model generates arguments (schema concern) but - # cannot supply request-time values like user (mapping concern). - tool_mapping = make_search_tool_mapping(user) - - if not previous_response_id: - response = client.responses.create( - input=[ - {"type": "message", "role": "user", "content": str(message)} - ], - **MODEL_DEFAULTS, - ) - else: - response = client.responses.create( + if previous_response_id: + initial_response = client.responses.create( input=[ {"type": "message", "role": "user", "content": str(message)} ], @@ -69,4 +39,15 @@ def run_assistant( **MODEL_DEFAULTS, ) - return handle_tool_calls_with_reasoning(response, client, MODEL_DEFAULTS, tool_mapping) + # search_documents needs the request user for document access control + return run_agentic_loop(initial_response, client, MODEL_DEFAULTS, TOOLS, user) + + initial_response = client.responses.create( + input=[ + {"type": "message", "role": "user", "content": str(message)} + ], + **MODEL_DEFAULTS, + ) + + # search_documents needs the request user for document access control + return run_agentic_loop(initial_response, client, MODEL_DEFAULTS, TOOLS, user) diff --git a/server/api/views/assistant/assistant_types.py b/server/api/views/assistant/assistant_types.py new file mode 100644 index 00000000..2f951672 --- /dev/null +++ b/server/api/views/assistant/assistant_types.py @@ -0,0 +1,92 @@ +from dataclasses import dataclass +from enum import Enum +from typing import Callable + +@dataclass(frozen=True) +class Tool: + """ + Instances are registered in tool_services.py's TOOLS list. + """ + name: str + description: str + parameters: dict + # Function we run: run(user, **arguments) -> str. + # Every tool takes the request `user` so the dispatch loop can call them uniformly; + # A tool that doesn't need it simply ignores it. + run: Callable + + # Schema that the model sees: Flattened Responses-API shape + def schema(self) -> dict: + return { + "type": "function", + "name": self.name, + "description": self.description, + "parameters": self.parameters, + } + + +class ToolCallStatus(str, Enum): + """ + Evaluate tool selection and distinguish between FAILED and UNREGISTERED + """ + + OK = "ok" + # Tool matched but raised an error + FAILED = "failed" + # No tool registered for the model's requested name: + # The model asked for a tool name we don't have + UNREGISTERED = "unregistered" + + +@dataclass(frozen=True) +class ToolCallExecution: + """ + A record of one tool call the model made + """ + + name: str + # `output` and `error` are disjoint by status + status: ToolCallStatus + # the query the model generated (the primary tool selection signal) + # None only when the model's argument JSON could not be parsed + arguments: dict | None = None + # the tool's result on success (retrieved content) + output: str | None = None + # the failure detail when status is not OK + error: str | None = None + +@dataclass(frozen=True) +class TokenUsage: + """ + Token usage for one responses.create call + + Every count is int | None, where None means unknown because response.usage + was missing, and 0 means no tokens were used + + There is no total_tokens field: it is input_tokens + output_tokens + """ + + input_tokens: int | None + # A subset of input_tokens + # Every iteration resends the context via previous_response_id + cached_input_tokens: int | None + output_tokens: int | None + # A subset of output_tokens + reasoning_output_tokens: int | None + +@dataclass(frozen=True) +class AgentResult: + """ + # Built by the agentic loop as a run proceeds, + # and read by eval_assistant.py to fill the result CSV + """ + + # The model's final text + output_text: str + # The id of the final response (for multi-turn continuity) + response_id: str + # The ordered ToolCallExecution records for every tool invocation across all loop iterations + tool_calls: list[ToolCallExecution] + # One per loop iteration. + # The eval sums these into its total_*_tokens columns + token_usages: list[TokenUsage] diff --git a/server/api/views/assistant/eval_assistant.py b/server/api/views/assistant/eval_assistant.py index b44a2174..cd2827e3 100644 --- a/server/api/views/assistant/eval_assistant.py +++ b/server/api/views/assistant/eval_assistant.py @@ -1,56 +1,59 @@ -#!/usr/bin/env -S uv run --script -# /// script -# requires-python = "==3.11.11" -# dependencies = [ -# "pandas==2.2.3", -# "openai", -# "django", -# ] -# /// - -# uv script (or plain Python) to generate results to CSV, run from the terminal -# Run from inside the container (working dir is /usr/src/server): -# docker compose exec backend python api/views/assistant/eval_assistant.py -# - +# Generates eval results to CSV. Run from inside the container: +# docker compose exec -e EVAL_BRANCH= backend python api/views/assistant/eval_assistant.py +# Writes to results/ next to this file, which the ./server bind mount surfaces on the host. import os import sys +import csv +import json import logging import datetime +from dataclasses import asdict +from time import perf_counter from concurrent.futures import ThreadPoolExecutor, as_completed -# Django setup must come before any imports that touch the ORM -# NOTE: from api/views/assistant/, "../../../../" resolves four levels up to -# /usr/src (not /usr/src/server, where balancer_backend lives). So this insert -# alone does not put the settings package on sys.path — running the script -# relies on the container already having /usr/src/server on PYTHONPATH. Sanity- -# check this the first time the eval is run for real; the path depth may need -# adjusting (e.g. "../../../"). -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../"))) +# Django setup must come before any imports that touch the ORM. +# Three levels up from api/views/assistant/ is /usr/src/server, where the balancer_backend settings package lives. +# Running a script file puts the *script's* directory on sys.path[0], not the working +# directory, and the image sets no PYTHONPATH — so without it django.setup() below +# raises ModuleNotFoundError on balancer_backend.settings. +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../"))) os.environ.setdefault("DJANGO_SETTINGS_MODULE", "balancer_backend.settings") import django django.setup() -from django.contrib.auth import get_user_model +from django.contrib.auth import get_user_model # noqa: E402 + +from api.views.assistant.assistant_services import run_assistant, MODEL_NAME # noqa: E402 +from api.views.assistant.assistant_types import ToolCallStatus +# Imported to warm the embedding model in main() before the worker pool starts — +# see the call site for why this process needs it and the web path does not. +from api.services.sentencetTransformer_model import TransformerModel # noqa: E402 -from api.views.assistant.assistant_services import run_assistant -# TODO: remove unused import or use INSTRUCTIONS to record an instructions_hash column -from api.views.assistant.assistant_prompts import INSTRUCTIONS logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s") logger = logging.getLogger(__name__) -# Read model and INSTRUCTIONS from the source file or add a lightweight config endpoint to the backend - -# Read model and INSTRUCTIONS from the source file -# INSTRUCTIONS is imported from assistant_prompts.py -# MODEL is read from assistant_services.py MODEL_DEFAULTS -# TODO: import a shared MODEL_NAME constant from assistant_services instead of hardcoding -MODEL = "gpt-5-nano" +FIELDNAMES = [ + "branch", + "question", + "error", + "response_output_text", + "duration_s", + # No total_tokens column: it is total_input_tokens + total_output_tokens + "total_input_tokens", + "total_cached_input_tokens", + "total_output_tokens", + "total_reasoning_output_tokens", + "total_tool_calls", + "total_tool_errors", + "tool_calls_json", + "token_usages_json", +] # Set of representative questions to evaluate the assistant + QUESTIONS = [ "What medications are recommended for bipolar depression?", "What are the risks of lithium for patients with kidney disease?", @@ -60,55 +63,73 @@ ] +def _total(iterations: list, field: str) -> int | None: + """Sum one token field across a turn's iterations, or None if any iteration's count is unknown + + Summed here rather than accumulated in the agentic loop + + A sum over only the known iterations would reach the CSV as a total count that is incomplete + """ + values = [getattr(iteration, field) for iteration in iterations] + if any(value is None for value in values): + return None + return sum(values) + + def run_one(question: str, user, branch: str) -> dict: """Run the assistant for a single question and return a result row. - Uses ThreadPoolExecutor (not asyncio.gather + await run_assistant) for concurrency. - - Concurrency approach comparison: - - ThreadPoolExecutor (this implementation): - - run_assistant stays sync — views.py and the WSGI web app are unaffected - - Each question runs in a thread pool worker, blocking on OpenAI + DB I/O - - Django DB safe when run via `docker compose exec backend python eval_assistant.py`: - this is a synchronous Django process context. Each ThreadPoolExecutor worker - is a real OS thread with its own threading.local() storage, so each thread - gets its own DB connection created lazily on first use. There is no shared - event loop thread, so connections cannot clash or bleed between questions. - The connection isolation concern only arises in ASGI contexts where multiple - coroutines share one thread and therefore one threading.local() connection — - which is not the case here. - - Runtime: bottlenecked by OpenAI rate limits, not thread overhead - - asyncio.gather + await run_assistant (alternative): - - run_assistant becomes async — requires async def post in views.py, - AsyncOpenAI client, and async handle_tool_calls_with_reasoning - - Django DB unsafe if get_closest_embeddings is called directly in an async - context without wrapping: get_closest_embeddings is a sync function that - hits the ORM, so calling it on the event loop thread blocks all other - coroutines until the DB responds. The fix is sync_to_async(get_closest_embeddings), - which runs it in a dedicated worker thread with its own threading.local() - connection. Bare await does not work at all — Django ORM querysets are not - awaitables and raise TypeError immediately. - - Under WSGI (manage.py runserver), async views run in a new event loop - per request — adds overhead to every web request for no benefit - - Cleaner call site in eval_assistant.py but wrong trade-off given WSGI + Uses ThreadPoolExecutor for concurrency. + """ + # Time the full run_assistant call here rather than inside it: run_one already + # owns the whole call, so wall-clock duration needs no plumbing through the production code path + start = perf_counter() try: - response_text, response_id = run_assistant(message=question, user=user) + result = run_assistant(message=question, user=user) + duration_s = perf_counter() - start return { "branch": branch, - "model": MODEL, "question": question, - "response_output_text": response_text, + # error is filled only when the whole turn failed and the except branch wrote the row "error": None, + "response_output_text": result.output_text, + "duration_s": duration_s, + "total_input_tokens": _total(result.token_usages, "input_tokens"), + "total_cached_input_tokens": _total(result.token_usages, "cached_input_tokens"), + "total_output_tokens": _total(result.token_usages, "output_tokens"), + "total_reasoning_output_tokens": _total(result.token_usages, "reasoning_output_tokens"), + # Flat summaries for scanning. total_tool_errors > 0 with error None means a tool + # failed inside the loop but the run still returned an answer. + "total_tool_calls": len(result.tool_calls), + # total_tool_errors counts tool calls whose status isn't OK (FAILED or UNREGISTERED). + # Those failures don't raise. The loop catches the exception, records it, and sends + # "Error executing function call: …" back to the model. + # The model then usually retries or writes a confident answer anyway + "total_tool_errors": sum(1 for c in result.tool_calls if c.status is not ToolCallStatus.OK), + # Per turn tool call details — status, the model's arguments (query), output/error + "tool_calls_json": json.dumps([asdict(c) for c in result.tool_calls]), + # Per iteration token usage detail the flat totals can't hold: which iteration caching engaged on + "token_usages_json": json.dumps([asdict(t) for t in result.token_usages]), } except Exception as e: + duration_s = perf_counter() - start logger.error(f"Error evaluating question '{question}': {e}") return { "branch": branch, - "model": MODEL, "question": question, - "response_output_text": None, "error": str(e), + "response_output_text": None, + "duration_s": duration_s, + # Tool calls and token usage before the error raised are not collected + "total_input_tokens": None, + "total_cached_input_tokens": None, + "total_output_tokens": None, + "total_reasoning_output_tokens": None, + "total_tool_calls": None, + "total_tool_errors": None, + "tool_calls_json": None, + "token_usages_json": None, } @@ -120,11 +141,15 @@ def main(): if not user: raise RuntimeError("No superuser found. Create one with manage.py createsuperuser.") - logger.info(f"Starting evaluation: branch={branch}, model={MODEL}, questions={len(QUESTIONS)}") + logger.info(f"Starting evaluation: branch={branch}, model={MODEL_NAME}, questions={len(QUESTIONS)}") - # ThreadPoolExecutor runs questions concurrently — see run_one docstring - # for trade-off discussion vs asyncio.gather + await run_assistant. - # max_workers=5 stays safely under OpenAI rate limits for gpt-5-nano. + # Load the embedding model before starting any workers + # TODO: Fix TransformerModel in its own commit — __new__ publishes _instance + # before .model loads, so concurrent callers get a half-built object + TransformerModel.get_instance() + + # ThreadPoolExecutor runs questions concurrently + # max_workers=5 stays safely under OpenAI rate limits for MODEL_NAME. results = [] with ThreadPoolExecutor(max_workers=5) as pool: futures = { @@ -134,18 +159,20 @@ def main(): for future in as_completed(futures): results.append(future.result()) - # Import pandas here, not at module top, so that importing this module (e.g. - # run_one from test_eval_assistant.py) does not require pandas. It is only - # needed for the CSV output below, when this script is run directly. - import pandas as pd - - df = pd.DataFrame(results) results_dir = os.path.join(os.path.dirname(__file__), "results") os.makedirs(results_dir, exist_ok=True) timestamp = datetime.datetime.utcnow().strftime("%Y%m%dT%H%M%S") output_path = os.path.join(results_dir, f"{branch}-{timestamp}.csv") - df.to_csv(output_path, index=False) + + # TODO: Write the system prompt next to the CSV ({branch}-{timestamp}.prompt.txt) so runs + # from before and after a prompt change can be told apart + + # pandas was never in the backend image's requirements.txt + with open(output_path, "w", newline="") as f: + writer = csv.DictWriter(f, fieldnames=FIELDNAMES) + writer.writeheader() + writer.writerows(results) logger.info(f"Results saved to {output_path}") diff --git a/server/api/views/assistant/search_tool.py b/server/api/views/assistant/search_tool.py new file mode 100644 index 00000000..cb6ebb4b --- /dev/null +++ b/server/api/views/assistant/search_tool.py @@ -0,0 +1,48 @@ +from api.services.embedding_services import get_closest_embeddings +from api.services.conversions_services import convert_uuids + + +def search_documents(query: str, user) -> str: + """ + Search through user's uploaded documents using semantic similarity. + + This function performs vector similarity search against the user's document corpus + and returns formatted results with context information for the LLM to use. + + Parameters + ---------- + query : str + The search query string + user : User + The authenticated user whose documents to search + + Returns + ------- + str + Formatted search results containing document excerpts with metadata, or a + message saying nothing matched. Matching nothing is a legitimate outcome, + not a failure, so it returns normally and the call is recorded as OK. + + Raises + ------ + Exception + If the embedding search fails. Deliberately not caught here. + + """ + + embeddings_results = get_closest_embeddings( + user=user, message_data=query.strip() + ) + embeddings_results = convert_uuids(embeddings_results) + + if not embeddings_results: + return "No relevant documents found for your query. Please try different search terms or upload documents first." + + # Format results with clear structure and metadata + # TODO: Drop `File: {file_id}` — the model sometimes cites this UUID instead of Name, which makes citations unparseable + prompt_texts = [ + f"[Document {i + 1} - File: {obj['file_id']}, Name: {obj['name']}, Page: {obj['page_number']}, Chunk: {obj['chunk_number']}, Similarity: {1 - obj['distance']:.3f}]\n{obj['text']}\n[End Document {i + 1}]" + for i, obj in enumerate(embeddings_results) + ] + + return "\n\n".join(prompt_texts) diff --git a/server/api/views/assistant/test_agentic_loop.py b/server/api/views/assistant/test_agentic_loop.py new file mode 100644 index 00000000..09f7b18c --- /dev/null +++ b/server/api/views/assistant/test_agentic_loop.py @@ -0,0 +1,104 @@ +# Unit tests for the assistant's highest-risk logic: the parts where a bug produces +# a confident answer with no evidence behind it, and nothing visibly fails. +# +# Deliberately not tested here (lower stakes, or fails loudly on the first real call): +# eval and token telemetry, previous_response_id handling in run_assistant, tool +# schemas and adapters, and the prompt text (model behavior, measured by the eval). +# +# TODO: Pin search_documents' result format once `File: {file_id}` is dropped (search_tool.py) + +import json +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from api.views.assistant.agentic_loop import handle_tool_calls, run_agentic_loop +from api.views.assistant.assistant_types import Tool, ToolCallStatus +from api.views.assistant.tool_services import SEARCH_TOOL + + +# Minimal stand-ins for OpenAI response objects: only the attributes the loop reads. +def call(name, call_id, **arguments): + return SimpleNamespace(type="function_call", name=name, call_id=call_id, arguments=json.dumps(arguments)) + + +def response(response_id, *items, text=None): + return SimpleNamespace(id=response_id, output=list(items), output_text=text) + + +# --- The loop grounds the answer in tool output -------------------------------------- +# +# If a tool's output is dropped, repeated, or paired with the wrong call_id, the model +# still writes a fluent answer, just without the retrieved evidence. The loop must also +# stop once the model stops calling tools; a reasoning item is not a tool call. + +def test_loop_feeds_each_tool_output_back_until_the_model_answers(): + tool = Tool("search_documents", "", {}, run=MagicMock(side_effect=["out1", "out2"])) + user = object() + client = MagicMock() + # A list side_effect runs out: if the loop failed to stop, the extra call would raise. + client.responses.create.side_effect = [ + response("r2", call("search_documents", "c2", query="q2")), + response("r3", SimpleNamespace(type="reasoning"), text="answer"), + ] + + result = run_agentic_loop( + response("r1", call("search_documents", "c1", query="q1")), client, {"instructions": "cite sources"}, [tool], user + ) + + # Each follow-up carries the previous iteration's output, paired with its call_id, + # and chains off that iteration's response id. It also resends the model settings: + # instructions are not carried over through previous_response_id, so without them + # the final answer would be written without the system prompt. + assert [ + (c.kwargs["previous_response_id"], c.kwargs["instructions"], c.kwargs["input"]) + for c in client.responses.create.call_args_list + ] == [ + ("r1", "cite sources", [{"type": "function_call_output", "call_id": "c1", "output": "out1"}]), + ("r2", "cite sources", [{"type": "function_call_output", "call_id": "c2", "output": "out2"}]), + ] + # Every dispatch gets the request user, which scopes document access. + assert [c.kwargs["user"] for c in tool.run.call_args_list] == [user, user] + assert (result.output_text, result.response_id) == ("answer", "r3") + + +# --- Tool failures are reported, never hidden ---------------------------------------- +# +# A failed call must reach the model as an error message (so it can retry or say it +# found nothing) and be recorded with its own status. A failure that raises out of the +# loop fails the user's request; one that looks like success turns into a confident, +# unsupported answer. + +def test_tool_failures_reach_the_model_and_are_recorded_by_kind(): + broken = Tool("search_documents", "", {}, run=MagicMock(side_effect=RuntimeError("db down"))) + + outputs, records = handle_tool_calls( + response("r1", call("made_up_tool", "c1"), call("search_documents", "c2", query="q")), [broken], user=None + ) + + assert [o["call_id"] for o in outputs] == ["c1", "c2"] + assert "made_up_tool" in outputs[0]["output"] and "db down" in outputs[1]["output"] + # Opposite diagnoses: the model invented a tool, versus our tool broke. + assert [r.status for r in records] == [ToolCallStatus.UNREGISTERED, ToolCallStatus.FAILED] + + +# --- Retrieval is scoped to the request user, and failure is not "no match" ---------- +# +# Runs the real SEARCH_TOOL with only the embedding search mocked. The request user +# must reach the search unchanged: it is what scopes document access, and passing the +# wrong one would silently return someone else's documents. And a search that breaks +# must be recorded as FAILED, not reported to the model as "nothing matched". + +@patch("api.views.assistant.search_tool.convert_uuids", side_effect=lambda rows: rows) +@patch("api.views.assistant.search_tool.get_closest_embeddings", side_effect=[RuntimeError("index down"), []]) +def test_search_uses_the_request_user_and_tells_failure_from_no_match(mock_search, _): + user = object() + + outputs, records = handle_tool_calls( + response("r1", call("search_documents", "c1", query="a"), call("search_documents", "c2", query="b")), + [SEARCH_TOOL], + user, + ) + + assert [c.kwargs["user"] for c in mock_search.call_args_list] == [user, user] + assert [r.status for r in records] == [ToolCallStatus.FAILED, ToolCallStatus.OK] + assert "No relevant documents found" in outputs[1]["output"] diff --git a/server/api/views/assistant/test_assistant_services.py b/server/api/views/assistant/test_assistant_services.py deleted file mode 100644 index 9d911920..00000000 --- a/server/api/views/assistant/test_assistant_services.py +++ /dev/null @@ -1,91 +0,0 @@ -# Tests for run_assistant (assistant_services.py): the orchestrator that wires the -# OpenAI client, the search tool mapping, and the agentic loop together. -# -# The OpenAI client and handle_tool_calls_with_reasoning are mocked, so these -# tests cover only logic run_assistant owns: how it builds the user input message, -# its decision to include vs. omit previous_response_id, and that it binds the -# request user into the search tool. No live OpenAI calls and no database. - -from unittest.mock import MagicMock, patch - - -def _make_terminal_response(output_text="Final answer.", response_id="resp-1"): - response = MagicMock() - response.output = [] - response.output_text = output_text - response.id = response_id - return response - -@patch("api.views.assistant.assistant_services.handle_tool_calls_with_reasoning") -@patch("api.views.assistant.assistant_services.OpenAI") -def test_run_assistant_sends_message_as_user_input(mock_openai_cls, mock_handle): - mock_client = MagicMock() - mock_openai_cls.return_value = mock_client - mock_client.responses.create.return_value = _make_terminal_response() - mock_handle.return_value = ("answer", "resp-1") - - from api.views.assistant.assistant_services import run_assistant - - run_assistant(message="Tell me about valproate.", user=MagicMock()) - - call_kwargs = mock_client.responses.create.call_args - input_messages = call_kwargs.kwargs.get("input") or call_kwargs.args[0] - assert any( - item.get("role") == "user" and "valproate" in item.get("content", "") - for item in input_messages - ) - - -@patch("api.views.assistant.assistant_services.handle_tool_calls_with_reasoning") -@patch("api.views.assistant.assistant_services.OpenAI") -def test_run_assistant_passes_previous_response_id(mock_openai_cls, mock_handle): - mock_client = MagicMock() - mock_openai_cls.return_value = mock_client - mock_client.responses.create.return_value = _make_terminal_response() - mock_handle.return_value = ("answer", "resp-2") - - from api.views.assistant.assistant_services import run_assistant - - run_assistant(message="More info.", user=MagicMock(), previous_response_id="resp-1") - - call_kwargs = mock_client.responses.create.call_args.kwargs - assert call_kwargs.get("previous_response_id") == "resp-1" - - -@patch("api.views.assistant.assistant_services.handle_tool_calls_with_reasoning") -@patch("api.views.assistant.assistant_services.OpenAI") -def test_run_assistant_omits_previous_response_id_when_none(mock_openai_cls, mock_handle): - mock_client = MagicMock() - mock_openai_cls.return_value = mock_client - mock_client.responses.create.return_value = _make_terminal_response() - mock_handle.return_value = ("answer", "resp-1") - - from api.views.assistant.assistant_services import run_assistant - - run_assistant(message="First message.", user=MagicMock(), previous_response_id=None) - - call_kwargs = mock_client.responses.create.call_args.kwargs - assert "previous_response_id" not in call_kwargs - - -@patch("api.views.assistant.tool_services.search_documents") -@patch("api.views.assistant.assistant_services.handle_tool_calls_with_reasoning") -@patch("api.views.assistant.assistant_services.OpenAI") -def test_run_assistant_binds_user_to_search_documents(mock_openai_cls, mock_handle, mock_search): - mock_client = MagicMock() - mock_openai_cls.return_value = mock_client - mock_client.responses.create.return_value = _make_terminal_response() - mock_handle.return_value = ("answer", "resp-1") - - from api.views.assistant.assistant_services import run_assistant - - user = MagicMock() - run_assistant(message="query", user=user) - - # Extract the tool_mapping passed to handle_tool_calls_with_reasoning - tool_mapping = mock_handle.call_args.kwargs.get("tool_mapping") or mock_handle.call_args.args[3] - bound_search = tool_mapping["search_documents"] - - # Calling the bound function should forward user to search_documents - bound_search(query="test query") - mock_search.assert_called_once_with("test query", user) diff --git a/server/api/views/assistant/test_eval_assistant.py b/server/api/views/assistant/test_eval_assistant.py deleted file mode 100644 index 5853d340..00000000 --- a/server/api/views/assistant/test_eval_assistant.py +++ /dev/null @@ -1,20 +0,0 @@ -# Tests for run_one (eval_assistant.py): the helper that runs the assistant for a -# single eval question and shapes the outcome into a result row. -# -# run_assistant is mocked, so this covers the logic run_one owns — specifically -# that a raising question is captured as an error row (error text recorded, -# response left None) instead of aborting the whole eval batch. - -from unittest.mock import MagicMock, patch - -from api.views.assistant.eval_assistant import run_one - -# TODO: add coverage for main()'s CSV output. - -@patch("api.views.assistant.eval_assistant.run_assistant", side_effect=Exception("boom")) -def test_run_one_captures_error(mock_run_assistant): - row = run_one("query", user=MagicMock(), branch="feature") - - assert row["branch"] == "feature" - assert row["response_output_text"] is None - assert "boom" in row["error"] diff --git a/server/api/views/assistant/test_tool_services.py b/server/api/views/assistant/test_tool_services.py deleted file mode 100644 index 86e57eed..00000000 --- a/server/api/views/assistant/test_tool_services.py +++ /dev/null @@ -1,218 +0,0 @@ -# Tests for tool_services.py: the retrieval tooling and the agentic reasoning loop. -# -# Covers the logic this module owns, with mocked tools (no DB, no OpenAI): -# - make_search_tool_mapping: the closure that binds the request user to -# search_documents, including per-call user independence. -# - invoke_functions_from_response: dispatching the model's function calls — -# the call/no-call branch, output shaping, and the unregistered-tool and -# tool-raises error paths. -# - handle_tool_calls_with_reasoning: the while-loop that keeps calling the -# model until it stops emitting tool calls, including loop continuity via -# previous_response_id. - -import json -from unittest.mock import MagicMock, patch - -# TODO: add coverage for search_documents itself (formatting of embeddings -# results, the empty-results message, and the exception path). No DB needed: -# search_documents only calls get_closest_embeddings and convert_uuids, so -# mocking those two (like the rest of the suite mocks collaborators) covers all -# three paths as fast, DB-free unit tests. - -from api.views.assistant.tool_services import ( - invoke_functions_from_response, - handle_tool_calls_with_reasoning, - make_search_tool_mapping, -) - - -# --------------------------------------------------------------------------- -# make_search_tool_mapping tests -# --------------------------------------------------------------------------- - -@patch("api.views.assistant.tool_services.search_documents") -def test_make_search_tool_mapping_bound_fn_forwards_user(mock_search): - mock_search.return_value = "results" - user = MagicMock() - mapping = make_search_tool_mapping(user) - - mapping["search_documents"](query="lithium") - - mock_search.assert_called_once_with("lithium", user) - - -@patch("api.views.assistant.tool_services.search_documents") -def test_make_search_tool_mapping_different_users_are_independent(mock_search): - # Each call to make_search_tool_mapping should capture its own user, - # so two mappings created with different users do not share state. - user_a = MagicMock() - user_b = MagicMock() - mapping_a = make_search_tool_mapping(user_a) - mapping_b = make_search_tool_mapping(user_b) - - mapping_a["search_documents"](query="q") - mapping_b["search_documents"](query="q") - - # bound_search calls search_documents(query, user) positionally, so each - # recorded call is (args, kwargs) == (("q", user), {}). - calls = mock_search.call_args_list - assert calls[0] == (("q", user_a), {}) - assert calls[1] == (("q", user_b), {}) - - -# --------------------------------------------------------------------------- -# invoke_functions_from_response tests -# --------------------------------------------------------------------------- - -def _make_function_call_item(name, arguments, call_id): - item = MagicMock() - item.type = "function_call" - item.name = name - item.arguments = json.dumps(arguments) - item.call_id = call_id - return item - - -def _make_reasoning_item(summary="reasoning summary"): - item = MagicMock() - item.type = "reasoning" - item.summary = summary - return item - - -def _make_response(output_items): - response = MagicMock() - response.output = output_items - return response - - -def test_invoke_returns_empty_list_when_no_function_calls(): - response = _make_response([_make_reasoning_item()]) - result = invoke_functions_from_response(response, tool_mapping={}) - assert result == [] - - -def test_invoke_calls_tool_and_returns_output(): - mock_tool = MagicMock(return_value="search result") - item = _make_function_call_item("search_documents", {"query": "lithium"}, "call-1") - response = _make_response([item]) - - result = invoke_functions_from_response( - response, tool_mapping={"search_documents": mock_tool} - ) - - mock_tool.assert_called_once_with(query="lithium") - assert result == [ - {"type": "function_call_output", "call_id": "call-1", "output": "search result"} - ] - - -def test_invoke_returns_error_message_when_tool_not_registered(): - item = _make_function_call_item("unknown_tool", {"query": "x"}, "call-2") - response = _make_response([item]) - - result = invoke_functions_from_response(response, tool_mapping={}) - - assert result[0]["call_id"] == "call-2" - assert "ERROR" in result[0]["output"] - - -def test_invoke_returns_error_message_when_tool_raises(): - mock_tool = MagicMock(side_effect=Exception("tool exploded")) - item = _make_function_call_item("search_documents", {"query": "x"}, "call-3") - response = _make_response([item]) - - result = invoke_functions_from_response( - response, tool_mapping={"search_documents": mock_tool} - ) - - assert "Error executing function call" in result[0]["output"] - - -def test_invoke_handles_multiple_function_calls(): - mock_tool = MagicMock(return_value="result") - items = [ - _make_function_call_item("search_documents", {"query": "q1"}, "call-4"), - _make_function_call_item("search_documents", {"query": "q2"}, "call-5"), - ] - response = _make_response(items) - - result = invoke_functions_from_response( - response, tool_mapping={"search_documents": mock_tool} - ) - - assert len(result) == 2 - assert mock_tool.call_count == 2 - - -# --------------------------------------------------------------------------- -# handle_tool_calls_with_reasoning tests -# --------------------------------------------------------------------------- - -def _make_terminal_response(output_text, response_id): - """A response with no function calls — terminates the loop.""" - response = MagicMock() - response.output = [] - response.output_text = output_text - response.id = response_id - return response - - -def _make_tool_call_response(response_id, query="lithium"): - """A response with one function call — continues the loop.""" - response = MagicMock() - response.output = [_make_function_call_item("search_documents", {"query": query}, "call-loop")] - response.id = response_id - return response - - -def test_handle_terminates_immediately_when_no_tool_calls(): - response = _make_terminal_response("Final answer.", "resp-1") - client = MagicMock() - - text, resp_id = handle_tool_calls_with_reasoning( - response, client, model_defaults={}, tool_mapping={} - ) - - assert text == "Final answer." - assert resp_id == "resp-1" - client.responses.create.assert_not_called() - - -def test_handle_calls_tool_then_terminates(): - mock_search = MagicMock(return_value="doc content") - first_response = _make_tool_call_response("resp-1") - second_response = _make_terminal_response("Final answer.", "resp-2") - - client = MagicMock() - client.responses.create.return_value = second_response - - text, resp_id = handle_tool_calls_with_reasoning( - first_response, - client, - model_defaults={}, - tool_mapping={"search_documents": mock_search}, - ) - - mock_search.assert_called_once_with(query="lithium") - assert text == "Final answer." - assert resp_id == "resp-2" - - -def test_handle_passes_previous_response_id_on_followup(): - mock_search = MagicMock(return_value="doc content") - first_response = _make_tool_call_response("resp-1") - second_response = _make_terminal_response("Done.", "resp-2") - - client = MagicMock() - client.responses.create.return_value = second_response - - handle_tool_calls_with_reasoning( - first_response, - client, - model_defaults={}, - tool_mapping={"search_documents": mock_search}, - ) - - call_kwargs = client.responses.create.call_args.kwargs - assert call_kwargs["previous_response_id"] == "resp-1" diff --git a/server/api/views/assistant/tool_services.py b/server/api/views/assistant/tool_services.py index 0fb96cef..eeecc6ea 100644 --- a/server/api/views/assistant/tool_services.py +++ b/server/api/views/assistant/tool_services.py @@ -1,214 +1,82 @@ -import json -import logging -from typing import Callable +from api.views.assistant.assistant_types import Tool +from api.views.assistant.search_tool import search_documents +from api.services.tools.database import ask_database -from ...services.embedding_services import get_closest_embeddings -from ...services.conversions_services import convert_uuids -logger = logging.getLogger(__name__) - -TOOL_DESCRIPTION = """ +SEARCH_TOOL = Tool( + name="search_documents", + description=""" Search the user's uploaded documents for information relevant to answering their question. Call this function when you need to find specific information from the user's documents to provide an accurate, citation-backed response. Always search before answering questions about document content. -""" - -TOOL_PROPERTY_DESCRIPTION = """ +""", + parameters={ + "type": "object", + "properties": { + "query": { + "type": "string", + "description": """ A specific search query to find relevant information in the user's documents. Use keywords, phrases, or questions related to what the user is asking about. Be specific rather than generic - use terms that would appear in the relevant documents. -""" - -# SEARCH_TOOLS_SCHEMA defines the search_documents tool for the OpenAI API. -# The model reads this schema to know what tools are available and what -# arguments to generate — it can only generate arguments declared here. -SEARCH_TOOLS_SCHEMA = [ - { - "type": "function", - "name": "search_documents", - "description": TOOL_DESCRIPTION, - "parameters": { - "type": "object", - "properties": { - "query": { - "type": "string", - "description": TOOL_PROPERTY_DESCRIPTION, - } - }, - "required": ["query"], +""", + } }, - } -] - - -# TODO: Add get_tools_schema() and make_tool_mapping(user) aggregation functions -# that combine all tool schemas and mappings so assistant_services.py never needs -# to change when a new tool is added — only tool_services.py does. - -def make_search_tool_mapping(user) -> dict[str, Callable]: - # make_search_tool_mapping binds user to search_documents at call time. - # user is a request-time value the model cannot generate, so it must be - # captured here and kept out of the schema. - """Return a tool mapping with search_documents bound to the given user. - - Parameters - ---------- - user : User - The Django user object used for document access control. - - Returns - ------- - dict[str, Callable] - Tool mapping ready to pass to invoke_functions_from_response. - """ - def bound_search(query: str) -> str: - return search_documents(query, user) - - return {"search_documents": bound_search} - - -def search_documents(query: str, user) -> str: - """ - Search through user's uploaded documents using semantic similarity. - - This function performs vector similarity search against the user's document corpus - and returns formatted results with context information for the LLM to use. - - Parameters - ---------- - query : str - The search query string - user : User - The authenticated user whose documents to search - - Returns - ------- - str - Formatted search results containing document excerpts with metadata - - Raises - ------ - Exception - If embedding search fails - """ - - try: - embeddings_results = get_closest_embeddings( - user=user, message_data=query.strip() - ) - embeddings_results = convert_uuids(embeddings_results) - - if not embeddings_results: - return "No relevant documents found for your query. Please try different search terms or upload documents first." - - # Format results with clear structure and metadata - prompt_texts = [ - f"[Document {i + 1} - File: {obj['file_id']}, Name: {obj['name']}, Page: {obj['page_number']}, Chunk: {obj['chunk_number']}, Similarity: {1 - obj['distance']:.3f}]\n{obj['text']}\n[End Document {i + 1}]" - for i, obj in enumerate(embeddings_results) - ] - - return "\n\n".join(prompt_texts) - - except Exception as e: - return f"Error searching documents: {str(e)}. Please try again if the issue persists." - - -def invoke_functions_from_response( - response, tool_mapping: dict[str, Callable] -) -> list[dict]: - """Extract all function calls from the response, look up the corresponding tool function(s) and execute them. - (This would be a good place to handle asynchroneous tool calls, or ones that take a while to execute.) - This returns a list of messages to be added to the conversation history. + "required": ["query"], + }, + + # search_documents needs the request user for document access control. + run=lambda user, query: search_documents(query, user), +) + +# The schema string describing the queryable medication table for ask_database's prompt. +# Kept in sync by hand with api.views.listMeds.models.Medication rather than deriving it from +# Django's Model._meta because the table is small and stable + +_MEDICATION_SCHEMA_STRING = "Table: api_medication\nColumns: name, benefits, risks" + + +# TODO: Rewrite the description as a directive like SEARCH_TOOL's — it documents SQL syntax instead +ASK_DATABASE_TOOL = Tool( + name="ask_database", + description=""" +Use this tool to answer questions about the medications in the Balancer database. +Medications are stored by their official generic names, not brand names, so convert +brand names to generic names first and match case-insensitively +(e.g. LOWER(name) = LOWER('lurasidone')). The input must be a single, fully-formed +SQL SELECT query. +""", + parameters={ + "type": "object", + "properties": { + "query": { + "type": "string", + "description": ( + "A plain-text SQL SELECT query answering the user's question, " + "written against this schema:\n" + f"{_MEDICATION_SCHEMA_STRING}" + ), + } + }, + "required": ["query"], + }, - Parameters - ---------- - response : OpenAI Response - The response object from OpenAI containing output items that may include function calls - tool_mapping : dict[str, Callable] - A dictionary mapping function names (as strings) to their corresponding Python functions. - Keys should match the function names defined in the tools schema. + # Reuses ask_database from services/tools. Its guards are substring checks, so they + # can be bypassed (UNION, a second statement), and it returns errors as strings, so + # a failed query would be recorded as OK. See the note at TOOLS. - Returns - ------- - list[dict] - List of function call output messages formatted for the OpenAI conversation. - Each message contains: - - type: "function_call_output" - - call_id: The unique identifier for the function call - - output: The result returned by the executed function (string or error message) - """ - - # Open AI Cookbook: Handling Function Calls with Reasoning Models - # https://cookbook.openai.com/examples/reasoning_function_calls - - intermediate_messages = [] - for response_item in response.output: - if response_item.type == "function_call": - target_tool = tool_mapping.get(response_item.name) - if target_tool: - try: - arguments = json.loads(response_item.arguments) - logger.info( - f"Invoking tool: {response_item.name} with arguments: {arguments}" - ) - tool_output = target_tool(**arguments) - logger.info(f"Tool {response_item.name} completed successfully") - except Exception as e: - msg = f"Error executing function call: {response_item.name}: {e}" - tool_output = msg - logger.error(msg, exc_info=True) - else: - msg = f"ERROR - No tool registered for function call: {response_item.name}" - tool_output = msg - logger.error(msg) - intermediate_messages.append( - { - "type": "function_call_output", - "call_id": response_item.call_id, - "output": tool_output, - } - ) - elif response_item.type == "reasoning": - logger.info(f"Reasoning step: {response_item.summary}") - return intermediate_messages + # ask_database queries the shared medication table, so it ignores the request user. + run=lambda user, query: ask_database(query), +) -def handle_tool_calls_with_reasoning( - response, client, model_defaults: dict, tool_mapping: dict[str, Callable] -) -> tuple[str, str]: - """Run the agentic loop until the model stops emitting function calls. - Parameters - ---------- - response : OpenAI Response - The initial response from the model. - client : OpenAI - The OpenAI client instance. - model_defaults : dict - Keyword arguments forwarded to every client.responses.create call. - tool_mapping : dict[str, Callable] - Maps function names to their implementations. +# Single source of truth for the assistant's tools. assistant_services builds the +# schema list the model sees with [tool.schema() for tool in TOOLS]; the agentic loop +# indexes this by name to dispatch calls. Register a new tool by appending it here. - Returns - ------- - tuple[str, str] - (final_response_output_text, final_response_id) - """ - # Open AI Cookbook: Handling Function Calls with Reasoning Models - # https://cookbook.openai.com/examples/reasoning_function_calls - while True: - # Mapping of the tool names we tell the model about and the functions that implement them - function_responses = invoke_functions_from_response(response, tool_mapping) - if len(function_responses) == 0: # We're done reasoning - logger.info("Reasoning completed") - final_response_output_text = response.output_text - final_response_id = response.id - logger.info(f"Final response: {final_response_output_text}") - return final_response_output_text, final_response_id - else: - logger.info("More reasoning required, continuing...") - response = client.responses.create( - input=function_responses, - previous_response_id=response.id, - **model_defaults, - ) +# ASK_DATABASE_TOOL is deliberately not registered: the endpoint is public (AllowAny), and +# ask_database can't safely run model-written SQL. Re-register it once it allows only a +# single statement, validates every table it reads, runs under a read-only database role, +# and raises on failure. +TOOLS = [SEARCH_TOOL] diff --git a/server/api/views/assistant/urls.py b/server/api/views/assistant/urls.py index 4c68f952..53467803 100644 --- a/server/api/views/assistant/urls.py +++ b/server/api/views/assistant/urls.py @@ -1,5 +1,5 @@ from django.urls import path -from .views import Assistant +from api.views.assistant.views import Assistant urlpatterns = [path("v1/api/assistant", Assistant.as_view(), name="assistant")] diff --git a/server/api/views/assistant/views.py b/server/api/views/assistant/views.py index 74bee8f6..5f988d86 100644 --- a/server/api/views/assistant/views.py +++ b/server/api/views/assistant/views.py @@ -9,7 +9,7 @@ from drf_spectacular.utils import extend_schema, inline_serializer from rest_framework import serializers as drf_serializers -from .assistant_services import run_assistant +from api.views.assistant.assistant_services import run_assistant logger = logging.getLogger(__name__) @@ -36,26 +36,21 @@ class Assistant(APIView): def post(self, request): try: user = request.user - - # TODO: validate message and return a 400 when it is omitted or blank. - # @extend_schema documents message as required, but that schema is not - # enforced at runtime, so a missing/empty message reaches run_assistant - # and becomes the literal string "None" (str(None)) in the model input — - # producing confusing model behavior. Add a 400 to the responses schema - # when implementing. + + # TODO: Missing/empty message reaches run_assistant and becomes the literal string "None" (str(None)) in the model input message = request.data.get("message", None) previous_response_id = request.data.get("previous_response_id", None) - final_response_output_text, final_response_id = run_assistant( - message=message, + result = run_assistant( user=user, + message=message, previous_response_id=previous_response_id, ) return Response( { - "response_output_text": final_response_output_text, - "final_response_id": final_response_id, + "response_output_text": result.output_text, + "final_response_id": result.response_id, }, status=status.HTTP_200_OK, ) diff --git a/server/balancer_backend/urls.py b/server/balancer_backend/urls.py index 55bd2032..999dee48 100644 --- a/server/balancer_backend/urls.py +++ b/server/balancer_backend/urls.py @@ -6,6 +6,9 @@ # Import TemplateView for rendering templates from django.views.generic import TemplateView import importlib # Import the importlib module for dynamic module importing +import os +from django.conf import settings +from django.http import HttpResponseNotFound from drf_spectacular.views import SpectacularAPIView, SpectacularSwaggerView, SpectacularRedocView @@ -58,10 +61,6 @@ path("api/redoc/", SpectacularRedocView.as_view(url_name="schema"), name="redoc"), ] -import os -from django.conf import settings -from django.http import HttpResponseNotFound - def spa_fallback(request): """Serve index.html for SPA routing when build is present; otherwise 404."""