| """Baseline inference runner for the B2B Support Triage OpenEnv benchmark.""" |
|
|
| from __future__ import annotations |
|
|
| import asyncio |
| import json |
| import os |
| import re |
| import textwrap |
| from typing import Any, Dict, List, Optional |
|
|
| from openai import OpenAI |
|
|
| from client import B2BSupportTriageEnv |
| from models import ActionType, B2BSupportPayload, B2BSupportTriageAction, B2BSupportTriageObservation |
|
|
| MODEL_NAME = os.getenv("MODEL_NAME") or "Qwen/Qwen2.5-72B-Instruct" |
| LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME") or os.getenv("IMAGE_NAME") or "b2b_support_triage_env-env:latest" |
|
|
| BENCHMARK = "b2b_support_triage_env" |
| TASKS = ["easy", "medium", "hard"] |
| TASK_SEEDS = {"easy": 101, "medium": 202, "hard": 303} |
| MAX_STEPS = 12 |
| MAX_TOKENS = 220 |
| TEMPERATURE = 0.0 |
| SUCCESS_SCORE_THRESHOLD = 0.80 |
|
|
| SYSTEM_PROMPT = textwrap.dedent( |
| """ |
| You are operating a B2B SaaS support triage environment. |
| Return ONLY compact JSON with this shape: |
| { |
| "action_type": "classify|set_priority|route|draft_reply|submit", |
| "ticket_id": "<ticket id or null for submit>", |
| "payload": { |
| "category": "...", |
| "priority": "...", |
| "route_queue": "...", |
| "sla_minutes": 120, |
| "escalate": true, |
| "reply_text": "..." |
| } |
| } |
| |
| Rules: |
| - Do not include markdown fences. |
| - Fill only payload keys needed for the chosen action_type. |
| - Keep action consistent with current plan and policy hints. |
| """ |
| ).strip() |
|
|
|
|
| def log_start(task: str, env: str, model: str) -> None: |
| print(f"[START] task={task} env={env} model={model}", flush=True) |
|
|
|
|
| def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None: |
| done_val = str(done).lower() |
| error_val = error if error else "null" |
| print( |
| f"[STEP] step={step} action={action} reward={reward:.2f} done={done_val} error={error_val}", |
| flush=True, |
| ) |
|
|
|
|
| def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None: |
| rewards_str = ",".join(f"{r:.2f}" for r in rewards) |
| print( |
| f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", |
| flush=True, |
| ) |
|
|
|
|
| def _extract_json_object(text: str) -> Dict[str, Any]: |
| candidate = text.strip() |
| if not candidate: |
| return {} |
|
|
| try: |
| return json.loads(candidate) |
| except json.JSONDecodeError: |
| pass |
|
|
| match = re.search(r"\{.*\}", candidate, re.DOTALL) |
| if not match: |
| return {} |
|
|
| try: |
| return json.loads(match.group(0)) |
| except json.JSONDecodeError: |
| return {} |
|
|
|
|
| def _deterministic_policy(obs: B2BSupportTriageObservation) -> B2BSupportTriageAction: |
| task = obs.task_id |
| ticket_id = obs.visible_ticket.ticket_id |
| decisions = obs.applied_decisions |
|
|
| targets = { |
| "easy": { |
| "category": "billing", |
| "priority": "medium", |
| "route_queue": "billing-general", |
| "sla_minutes": 480, |
| "escalate": False, |
| }, |
| "medium": { |
| "category": "billing", |
| "priority": "high", |
| "route_queue": "billing-l2", |
| "sla_minutes": 120, |
| "escalate": False, |
| }, |
| "hard": { |
| "category": "security", |
| "priority": "urgent", |
| "route_queue": "security-incident-response", |
| "sla_minutes": 120, |
| "escalate": True, |
| }, |
| } |
|
|
| target = targets[task] |
|
|
| if "category" not in decisions: |
| return B2BSupportTriageAction( |
| action_type=ActionType.CLASSIFY, |
| ticket_id=ticket_id, |
| payload=B2BSupportPayload(category=target["category"]), |
| ) |
|
|
| if "priority" not in decisions: |
| return B2BSupportTriageAction( |
| action_type=ActionType.SET_PRIORITY, |
| ticket_id=ticket_id, |
| payload=B2BSupportPayload(priority=target["priority"]), |
| ) |
|
|
| if "route_queue" not in decisions or "sla_minutes" not in decisions: |
| return B2BSupportTriageAction( |
| action_type=ActionType.ROUTE, |
| ticket_id=ticket_id, |
| payload=B2BSupportPayload( |
| route_queue=target["route_queue"], |
| sla_minutes=target["sla_minutes"], |
| escalate=target["escalate"] if task == "hard" else None, |
| ), |
| ) |
|
|
| if task == "hard" and "reply_text" not in decisions: |
| reply_text = ( |
| "We have escalated this to our security team. " |
| "The incident is escalated and under active investigation. " |
| "Please reset your API key immediately; we will share an update within 2 hours." |
| ) |
| return B2BSupportTriageAction( |
| action_type=ActionType.DRAFT_REPLY, |
| ticket_id=ticket_id, |
| payload=B2BSupportPayload(reply_text=reply_text), |
| ) |
|
|
| return B2BSupportTriageAction(action_type=ActionType.SUBMIT, ticket_id=None, payload=B2BSupportPayload()) |
|
|
|
|
| def _coerce_model_action(raw: Dict[str, Any], obs: B2BSupportTriageObservation) -> Optional[B2BSupportTriageAction]: |
| if not raw: |
| return None |
|
|
| try: |
| action_type = ActionType(raw.get("action_type", "")) |
| except Exception: |
| return None |
|
|
| payload = raw.get("payload") or {} |
| ticket_id = raw.get("ticket_id") |
|
|
| if action_type != ActionType.SUBMIT and not ticket_id: |
| ticket_id = obs.visible_ticket.ticket_id |
|
|
| try: |
| return B2BSupportTriageAction( |
| action_type=action_type, |
| ticket_id=ticket_id, |
| payload=B2BSupportPayload( |
| category=payload.get("category"), |
| priority=payload.get("priority"), |
| route_queue=payload.get("route_queue"), |
| sla_minutes=payload.get("sla_minutes"), |
| escalate=payload.get("escalate"), |
| reply_text=payload.get("reply_text"), |
| ), |
| ) |
| except Exception: |
| return None |
|
|
|
|
| def _action_to_string(action: B2BSupportTriageAction) -> str: |
| payload = action.payload.model_dump(exclude_none=True) |
| compact = {"action_type": action.action_type.value, "ticket_id": action.ticket_id, "payload": payload} |
| return json.dumps(compact, separators=(",", ":"), ensure_ascii=True) |
|
|
|
|
| def _build_user_prompt(step: int, obs: B2BSupportTriageObservation, history: List[str]) -> str: |
| return textwrap.dedent( |
| f""" |
| Step: {step} |
| Task: {obs.task_id} |
| Ticket ID: {obs.visible_ticket.ticket_id} |
| Subject: {obs.visible_ticket.subject} |
| Body: {obs.visible_ticket.body} |
| Current decisions: {json.dumps(obs.applied_decisions, ensure_ascii=True)} |
| Current plan: {json.dumps(obs.current_plan, ensure_ascii=True)} |
| Last action error: {obs.last_action_error} |
| Progress score: {obs.progress_score:.3f} |
| Last 4 history lines: {history[-4:] if history else []} |
| |
| Output one JSON action object only. |
| """ |
| ).strip() |
|
|
|
|
| def _call_model_action(client: OpenAI, step: int, obs: B2BSupportTriageObservation, history: List[str]) -> Dict[str, Any]: |
| prompt = _build_user_prompt(step, obs, history) |
| completion = client.chat.completions.create( |
| model=MODEL_NAME, |
| messages=[ |
| {"role": "system", "content": SYSTEM_PROMPT}, |
| {"role": "user", "content": prompt}, |
| ], |
| temperature=TEMPERATURE, |
| max_tokens=MAX_TOKENS, |
| stream=False, |
| ) |
| content = (completion.choices[0].message.content or "").strip() |
| return _extract_json_object(content) |
|
|
|
|
| def _touch_proxy(client: OpenAI) -> None: |
| """Force at least one LiteLLM proxy request even if env execution fails early.""" |
| try: |
| _ = client.models.list() |
| return |
| except Exception: |
| pass |
|
|
| |
| try: |
| _ = client.chat.completions.create( |
| model=MODEL_NAME, |
| messages=[{"role": "user", "content": "Reply with JSON: {}"}], |
| temperature=0.0, |
| max_tokens=4, |
| stream=False, |
| ) |
| except Exception: |
| pass |
|
|
|
|
| async def run_single_task(client: OpenAI, task_name: str, seed: int) -> float: |
| rewards: List[float] = [] |
| history: List[str] = [] |
| steps_taken = 0 |
| final_score = 0.0 |
| success = False |
|
|
| log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME) |
|
|
| env: Optional[B2BSupportTriageEnv] = None |
| try: |
| env = await B2BSupportTriageEnv.from_docker_image(LOCAL_IMAGE_NAME) |
| result = await env.reset(task_id=task_name, seed=seed) |
|
|
| for step in range(1, MAX_STEPS + 1): |
| if result.done: |
| break |
|
|
| obs = result.observation |
| deterministic = _deterministic_policy(obs) |
|
|
| model_raw: Dict[str, Any] = {} |
| try: |
| model_raw = _call_model_action(client, step, obs, history) |
| except Exception: |
| model_raw = {} |
|
|
| model_action = _coerce_model_action(model_raw, obs) |
| action = model_action if (model_action and model_action.action_type == deterministic.action_type) else deterministic |
|
|
| result = await env.step(action) |
| reward = float(result.reward or 0.0) |
| done = bool(result.done) |
| error = result.observation.last_action_error |
|
|
| rewards.append(reward) |
| steps_taken = step |
| history.append(f"step={step} action={action.action_type.value} reward={reward:.2f}") |
|
|
| log_step(step=step, action=_action_to_string(action), reward=reward, done=done, error=error) |
|
|
| if done: |
| break |
|
|
| final_score = float(result.observation.progress_score) if steps_taken > 0 else 0.0 |
| success = final_score >= SUCCESS_SCORE_THRESHOLD |
|
|
| except Exception: |
| success = False |
|
|
| finally: |
| if env is not None: |
| try: |
| await env.close() |
| except Exception: |
| pass |
|
|
| log_end(success=success, steps=steps_taken, score=final_score, rewards=rewards) |
|
|
| return final_score |
|
|
|
|
| async def main() -> None: |
| api_key = os.getenv("API_KEY") or os.getenv("HF_TOKEN") |
| if not api_key: |
| raise RuntimeError("Missing API key: set API_KEY or HF_TOKEN") |
|
|
| client = OpenAI( |
| base_url=os.environ["API_BASE_URL"], |
| api_key=api_key, |
| ) |
| _touch_proxy(client) |
|
|
| scores: List[float] = [] |
| for task_name in TASKS: |
| score = await run_single_task(client, task_name, TASK_SEEDS[task_name]) |
| scores.append(score) |
|
|
| aggregate = sum(scores) / len(scores) if scores else 0.0 |
| print(f"Baseline aggregate score: {aggregate:.3f}", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| asyncio.run(main()) |
|
|