import json from typing import AsyncGenerator from openai import AsyncOpenAI from config import LLM_BASE_URL, LLM_MODEL, OPENAI_API_KEY from tools import TOOLS, run_tool _SYSTEM_INJECT = ( "You have access to a live notepad. " "Use write_to_notepad whenever the user asks you to note, write down, remember, save, " "or record anything — always confirm what you wrote after calling the tool. " "Use read_notepad to recall saved notes. " "Use clear_notepad only when the user explicitly asks to erase all notes. " "Use get_datetime for the current date or time.\n\n" "You also act as a clinical assistant for blood test results. " "Whenever the user asks about their blood test, lab results, deficiencies, or what " "supplements they should take, first call get_blood_report_summary to see which " "markers are low, high, or normal. For every abnormal (low or high) marker, call " "search_supplements with a short query describing that marker (e.g. 'low vitamin D') " "to get supplement recommendations, then explain the deficiency and recommendation in " "plain spoken language — no markdown, bullet points, asterisks, or headers, since this " "is read aloud by a speaking avatar. If the search_supplements result includes " "interaction cautions (e.g. a medication, pregnancy, or condition to be careful with), " "briefly mention that caution too, in plain language — this matters most for markers " "like kidney function, cholesterol/lipids, and blood sugar where the wrong supplement " "can be risky. Always end any discussion of blood results or supplements with a " "reminder that this is not a medical diagnosis and the user should consult their " "doctor before starting any supplement or changing their dose." ) _client: AsyncOpenAI | None = None def _get_client() -> AsyncOpenAI: global _client if _client is None: _client = AsyncOpenAI(api_key=OPENAI_API_KEY, base_url=LLM_BASE_URL) return _client def _build_history(messages: list) -> list: """Inject tool guidance into the system message.""" msgs = list(messages) if msgs and msgs[0].get("role") == "system": msgs[0] = {**msgs[0], "content": _SYSTEM_INJECT + "\n\n" + msgs[0]["content"]} else: msgs.insert(0, {"role": "system", "content": _SYSTEM_INJECT}) return msgs async def run_agent(messages: list, model: str | None = None) -> str: """Run the full tool-calling loop and return the final text response.""" history = _build_history(messages) m = model or LLM_MODEL print(f"[AGENT → {LLM_BASE_URL}] model={m} msgs={len(history)}", flush=True) while True: resp = await _get_client().chat.completions.create( model=m, messages=history, tools=TOOLS, tool_choice="auto", store=True, ) msg = resp.choices[0].message if not msg.tool_calls: print(f"[AGENT] ✓ final reply: {(msg.content or '')!r:.100}", flush=True) return msg.content or "" history.append(msg.model_dump(exclude_none=True)) for tc in msg.tool_calls: args = json.loads(tc.function.arguments or "{}") print(f"[AGENT] tool call: {tc.function.name}({args})", flush=True) result = await run_tool(tc.function.name, args) history.append({"role": "tool", "tool_call_id": tc.id, "content": result}) async def stream_agent(messages: list, model: str | None = None) -> AsyncGenerator[str, None]: """ Stream tokens as soon as the model produces them. Every pass is a single streaming call — if the model asks for tool calls we resolve them and loop with another streaming call; if it answers directly, those tokens are yielded immediately with no extra non-streamed round trip first. """ history = _build_history(messages) m = model or LLM_MODEL while True: stream = await _get_client().chat.completions.create( model=m, messages=history, tools=TOOLS, tool_choice="auto", stream=True, store=True, ) tool_calls: dict[int, dict] = {} async for chunk in stream: delta = chunk.choices[0].delta if delta.content: yield delta.content for tc_delta in (delta.tool_calls or []): slot = tool_calls.setdefault( tc_delta.index, {"id": "", "name": "", "arguments": ""} ) if tc_delta.id: slot["id"] = tc_delta.id if tc_delta.function: if tc_delta.function.name: slot["name"] += tc_delta.function.name if tc_delta.function.arguments: slot["arguments"] += tc_delta.function.arguments if not tool_calls: return # final answer was already streamed above ordered = [tool_calls[i] for i in sorted(tool_calls)] history.append({ "role": "assistant", "content": None, "tool_calls": [ { "id": tc["id"], "type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}, } for tc in ordered ], }) for tc in ordered: args = json.loads(tc["arguments"] or "{}") print(f"[AGENT] tool call: {tc['name']}({args})", flush=True) result = await run_tool(tc["name"], args) history.append({"role": "tool", "tool_call_id": tc["id"], "content": result})