winniekariuki's picture
Phase 3 complete: working chatbot with policy gate, tests, Gradio UI
f7e2a40
Raw
History Blame Contribute Delete
18.7 kB
"""OpenAI tool-calling loop with mandatory policy gating before MCP I/O."""
from __future__ import annotations
import json
import os
import re
import time
import uuid
from dataclasses import dataclass
from typing import Any
from openai import APIError, AsyncOpenAI
from src.logger import log_event, summarize_tool_args
from src.mcp_client import MCPAuthError, MCPNotFoundError, MCPTimeoutError, MCPTransportError, MeridianMCPClient
from src.policy import Allow, Deny, RequireConfirmation, authorize, inject_sensitive_tool_args, sync_pending_order_snapshot
from src.prompts import build_system_prompt
from src.schemas import LLM_TOOL_NAMES, TOOL_INPUT_MODELS
from src.session import Message, Session
_FALLBACK_MODEL = "gpt-4o-mini"
_MAX_TOOL_ITERATIONS = 5
MCP_TOOL_WHITELIST: frozenset[str] = frozenset(TOOL_INPUT_MODELS.keys())
_CONFIRM_RE = re.compile(
r"^(yes|yep|yeah|confirm|confirmed|place\s+(?:the\s+)?order|go\s+ahead|"
r"do\s+it|ok(?:ay)?|sounds\s+good|please\s+do)\b",
re.I | re.S,
)
_REJECT_RE = re.compile(r"^(no|nope|nah|cancel|stop|wait|not\s+now)\b", re.I | re.S)
_SKU_RE = re.compile(r"\b([A-Z]{2,12}-\d{2,8})\b")
@dataclass
class ToolCallRecord:
tool_name: str
args_summary: str
decision: str | None = None
latency_ms: int | None = None
status: str | None = None
@dataclass
class TurnResult:
assistant_text: str
tool_calls_made: list[ToolCallRecord]
confirmation_requested: bool
error: str | None = None
def apply_user_confirmation_intent(session: Session, user_message: str) -> None:
"""Deterministic confirmation handling — never delegate this to the LLM."""
text = user_message.strip()
if not text:
return
pending = session.pending_confirmation
if not pending:
return
if _REJECT_RE.match(text):
session.pending_confirmation = None
return
if _CONFIRM_RE.match(text):
pending.confirmed = True
def _extract_price(text: str) -> str | None:
m = re.search(r"\$?\s*(\d+(?:\.\d{1,2})?)", text)
return m.group(1) if m else None
def update_facts_from_mcp_text(tool_name: str, text: str, session: Session) -> None:
"""Best-effort fact extraction — MCP returns prose, not JSON."""
if tool_name in {"search_products", "list_products"}:
results: list[dict[str, str | None]] = []
for line in text.splitlines():
sku_m = _SKU_RE.search(line)
if not sku_m:
continue
results.append(
{
"sku": sku_m.group(1),
"name": line.strip()[:160],
"price": _extract_price(line),
}
)
if len(results) >= 5:
break
if results:
session.facts["last_search_results"] = results
return
if tool_name == "get_product":
sku_m = _SKU_RE.search(text)
top = text.splitlines()[0][:220] if text else ""
session.facts["last_viewed_product"] = {
"sku": sku_m.group(1) if sku_m else None,
"name": top,
"price": _extract_price(text),
}
return
if tool_name == "list_orders":
orders: list[dict[str, str | None]] = []
uuid_re = re.compile(
r"[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}",
re.I,
)
for line in text.splitlines():
oid_m = uuid_re.search(line)
if not oid_m:
continue
status: str | None = None
for st in ("draft", "submitted", "approved", "fulfilled", "cancelled"):
if st in line.lower():
status = st
break
orders.append(
{
"order_id": oid_m.group(0),
"status": status,
"total": _extract_price(line),
}
)
if len(orders) >= 5:
break
if orders:
session.facts["last_orders"] = orders
return
if tool_name == "create_order":
session.pending_confirmation = None
uuid_re = re.compile(
r"[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}",
re.I,
)
oid_m = uuid_re.search(text)
entry = {"order_id": oid_m.group(0) if oid_m else None, "snippet": text[:400]}
recent = session.facts.get("recent_orders")
buf: list[Any] = recent if isinstance(recent, list) else []
session.facts["recent_orders"] = [entry] + buf[:4]
def _mcp_tool_to_openai(tool: Any) -> dict[str, Any]:
return {
"type": "function",
"function": {
"name": tool.name,
"description": (tool.description or "")[:4096],
"parameters": tool.inputSchema,
},
}
def _history_to_messages(session: Session) -> list[dict[str, Any]]:
"""Prior turns only — the active user utterance is appended by handle_turn."""
msgs: list[dict[str, Any]] = []
hist = list(session.history)
if len(hist) < 2:
return msgs
for m in hist[:-1]:
if m.role not in {"user", "assistant"}:
continue
msgs.append({"role": m.role, "content": m.content})
return msgs
class ChatEngine:
def __init__(
self,
mcp_client: MeridianMCPClient,
session: Session,
openai_client: AsyncOpenAI,
*,
session_id: str | None = None,
model: str | None = None,
) -> None:
self._mcp = mcp_client
self._session = session
self._openai = openai_client
self._session_id = session_id or str(uuid.uuid4())
self._model = model or os.environ.get("OPENAI_MODEL", "").strip() or _FALLBACK_MODEL
self._tools_cache: list[dict[str, Any]] | None = None
@property
def session_id(self) -> str:
return self._session_id
async def _openai_tools(self) -> list[dict[str, Any]]:
if self._tools_cache is None:
remote = await self._mcp.list_tools()
self._tools_cache = [_mcp_tool_to_openai(t) for t in remote if t.name in LLM_TOOL_NAMES]
return self._tools_cache
async def handle_turn(self, user_message: str) -> TurnResult:
apply_user_confirmation_intent(self._session, user_message)
self._session.add_message(Message(role="user", content=user_message))
hist = list(self._session.history)
if not hist or hist[-1].role != "user":
return TurnResult(
"I'm missing internal chat state. Please restart the session.",
[],
False,
error="history_invariant_broken",
)
prior_user_text = hist[-1].content
api_messages: list[dict[str, Any]] = _history_to_messages(self._session) + [
{"role": "user", "content": prior_user_text}
]
tool_records: list[ToolCallRecord] = []
confirmation_requested = False
tools = await self._openai_tools()
for iteration in range(_MAX_TOOL_ITERATIONS):
system_prompt = build_system_prompt(self._session)
t_llm = time.perf_counter()
try:
completion = await self._openai.chat.completions.create(
model=self._model,
max_tokens=1024,
messages=[{"role": "system", "content": system_prompt}, *api_messages],
tools=tools,
tool_choice="auto",
)
except APIError as e:
log_event(
"error",
session_id=self._session_id,
tool=None,
decision=None,
deny_reason=str(e),
latency_ms=int((time.perf_counter() - t_llm) * 1000),
status="error",
args_summary=None,
)
msg = "I'm experiencing a technical issue. Please try again."
self._session.add_message(Message(role="assistant", content=msg))
return TurnResult(msg, tool_records, confirmation_requested, error=str(e))
choice = completion.choices[0]
finish = choice.finish_reason
msg = choice.message
log_event(
"llm_call",
session_id=self._session_id,
tool=None,
decision=None,
deny_reason=None,
latency_ms=int((time.perf_counter() - t_llm) * 1000),
status="success",
args_summary=f"finish_reason={finish};iteration={iteration + 1}",
)
tool_calls = msg.tool_calls or []
if finish != "tool_calls" or not tool_calls:
text = (msg.content or "").strip() or (
"I'm not sure how to help with that — could you rephrase?"
)
self._session.add_message(Message(role="assistant", content=text))
return TurnResult(text, tool_records, confirmation_requested, None)
assistant_msg: dict[str, Any] = {
"role": "assistant",
"content": msg.content or "",
"tool_calls": [
{
"id": tc.id,
"type": "function",
"function": {
"name": tc.function.name,
"arguments": tc.function.arguments or "{}",
},
}
for tc in tool_calls
],
}
api_messages.append(assistant_msg)
for tc in tool_calls:
tool_name = tc.function.name
tool_id = tc.id
raw_args = tc.function.arguments or "{}"
try:
parsed = json.loads(raw_args) if raw_args.strip() else {}
except json.JSONDecodeError:
parsed = {}
if not isinstance(parsed, dict):
summary = summarize_tool_args(tool_name, {})
tool_records.append(
ToolCallRecord(
tool_name,
summary,
decision=None,
latency_ms=None,
status="bad_args",
)
)
log_event(
"error",
session_id=self._session_id,
tool=tool_name,
decision=None,
deny_reason="invalid_tool_arguments_json",
latency_ms=None,
status="error",
args_summary=summary,
)
api_messages.append(
{
"role": "tool",
"tool_call_id": tool_id,
"content": "Tool arguments must be a JSON object; please fix and retry.",
}
)
continue
args = dict(parsed)
summary = summarize_tool_args(tool_name, args)
t0 = time.perf_counter()
if tool_name not in MCP_TOOL_WHITELIST:
lat = int((time.perf_counter() - t0) * 1000)
tool_records.append(
ToolCallRecord(tool_name, summary, decision="deny", latency_ms=lat, status="blocked")
)
log_event(
"policy_decision",
session_id=self._session_id,
tool=tool_name,
decision="deny",
deny_reason="Unknown tool",
latency_ms=lat,
status=None,
args_summary=summary,
)
api_messages.append(
{
"role": "tool",
"tool_call_id": tool_id,
"content": "Unknown tool — pick from the registered Meridian MCP tools only.",
}
)
continue
decision = authorize(tool_name, args, self._session)
if isinstance(decision, Deny):
lat = int((time.perf_counter() - t0) * 1000)
tool_records.append(
ToolCallRecord(
tool_name,
summary,
decision="deny",
latency_ms=lat,
status="denied",
)
)
log_event(
"policy_decision",
session_id=self._session_id,
tool=tool_name,
decision="deny",
deny_reason=decision.reason,
latency_ms=lat,
status=None,
args_summary=summary,
)
api_messages.append(
{"role": "tool", "tool_call_id": tool_id, "content": decision.reason}
)
continue
if isinstance(decision, RequireConfirmation):
lat = int((time.perf_counter() - t0) * 1000)
confirmation_requested = True
tool_records.append(
ToolCallRecord(
tool_name,
summary,
decision="require_confirmation",
latency_ms=lat,
status="pending",
)
)
log_event(
"policy_decision",
session_id=self._session_id,
tool=tool_name,
decision="require_confirmation",
deny_reason=None,
latency_ms=lat,
status=None,
args_summary=summary,
)
if tool_name == "create_order" and self._session.pending_confirmation is None:
canonical = inject_sensitive_tool_args(
tool_name,
args,
self._session,
)
sync_pending_order_snapshot(self._session, canonical)
api_messages.append(
{"role": "tool", "tool_call_id": tool_id, "content": decision.prompt}
)
continue
merged = inject_sensitive_tool_args(tool_name, args, self._session)
log_event(
"policy_decision",
session_id=self._session_id,
tool=tool_name,
decision="allow",
deny_reason=None,
latency_ms=int((time.perf_counter() - t0) * 1000),
status=None,
args_summary=summary,
)
mcp_t0 = time.perf_counter()
try:
payload = await self._mcp.call_tool(tool_name, merged)
mcp_lat = int((time.perf_counter() - mcp_t0) * 1000)
update_facts_from_mcp_text(tool_name, payload, self._session)
tool_records.append(
ToolCallRecord(
tool_name,
summarize_tool_args(tool_name, merged),
decision="allow",
latency_ms=mcp_lat,
status="success",
)
)
log_event(
"tool_call",
session_id=self._session_id,
tool=tool_name,
decision="allow",
deny_reason=None,
latency_ms=mcp_lat,
status="success",
args_summary=summarize_tool_args(tool_name, merged),
)
api_messages.append(
{"role": "tool", "tool_call_id": tool_id, "content": payload}
)
except (MCPTransportError, MCPAuthError, MCPNotFoundError, MCPTimeoutError) as e:
mcp_lat = int((time.perf_counter() - mcp_t0) * 1000)
tool_records.append(
ToolCallRecord(
tool_name,
summarize_tool_args(tool_name, merged),
decision="allow",
latency_ms=mcp_lat,
status="error",
)
)
log_event(
"error",
session_id=self._session_id,
tool=tool_name,
decision="allow",
deny_reason=str(e),
latency_ms=mcp_lat,
status="error",
args_summary=summarize_tool_args(tool_name, merged),
)
api_messages.append(
{
"role": "tool",
"tool_call_id": tool_id,
"content": "I couldn't reach that information right now. Please try again shortly.",
}
)
fallback = "I'm having trouble completing that request. Could you rephrase?"
self._session.add_message(Message(role="assistant", content=fallback))
log_event(
"error",
session_id=self._session_id,
tool=None,
decision=None,
deny_reason="tool_iteration_cap",
latency_ms=None,
status="error",
args_summary=None,
)
return TurnResult(fallback, tool_records, confirmation_requested, error="tool_iteration_cap")