StarTripper/ticket_ordering
0
1import os2import asyncio3import textwrap4import numpy as np5from openai import OpenAI6from typing import List, Optional, Dict, Any7 8from client import TicketOrderingEnv9from problem_generator import GenerationDifficulty10from models import TicketOrderingAction, TicketOrderingObservation11 12 13backup_rng = np.random.default_rng(42)14 15 16ENV_BASE = "https://startripper-openenv-ticket-ordering-env.hf.space"17 18API_KEY = os.getenv("HF_TOKEN") or os.getenv("API_KEY")19API_BASE_URL = os.getenv("API_BASE_URL", "https://api.groq.com/openai/v1")20MODEL_NAME = os.getenv("MODEL_NAME") or "llama-3.1-8b-instant"21MAX_STEPS = 1522TEMPERATURE = 0.0 # Kind of makes these models TOO deterministic / repeat things, oh well, rules are rules.23MAX_TOKENS = 30024SUCCESS_SCORE_THRESHOLD = 0.7525 26SYSTEM_PROMPT = textwrap.dedent(27 """28 You are a ticket prioritization agent.29 30 Given:31 - ordering criteria32 - reference tickets33 - a candidate ticket34 - heuristics from n of the most assigned and n of the least assigned tickets35 - total number of tickets36 - iterations completed so far37 38 Your job:39 - Assign a priority score to the candidate (relative to the references, must NOT be neutral / 0.0) (float, higher = more important)40 - Write a short summary for the candidate (<=32 chars) (Part of that ticket's heuristic)41 - Select next reference ticket ids (must be one of the keys from the heuristics)42 - Select next candidate ticket id (must be one of the keys from the heuristics)43 - Decide whether to end ordering (there is no need to end after iterations completed = total tickets since cross comparing tickets still may be valuable)44 45 Respond with no fancy formatting, backticks or anything else, respond ONLY with PURE valid JSON in this format:46 {47 "candidate_priority": float,48 "candidate_summary": string,49 "next_reference_ids": [int],50 "next_candidate_id": int,51 "end_ordering": bool52 }53 """54).strip()55 56 57def log_start(task: str, env: str, model: str) -> None:58 print(f"[START] task={task} env={env} model={model}", flush=True)59 60 61def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:62 error_val = error if error else "null"63 done_val = str(done).lower()64 print(65 f"[STEP] step={step} action={action} reward={reward:.2f} done={done_val} error={error_val}",66 flush=True,67 )68 69 70def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:71 rewards_str = ",".join(f"{r:.2f}" for r in rewards)72 print(f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", flush=True)73 74 75def serialize_ticket(ticket: Any) -> Dict[str, Any]:76 return {77 "id": ticket.id,78 "thread": [{"user": c.user, "content": c.content} for c in ticket.thread],79 "heuristic": {80 "priority": ticket.heuristic.priority,81 "summary": ticket.heuristic.summary,82 "times_assigned": ticket.heuristic.times_assigned,83 },84 }85 86 87def build_user_prompt(obs: Any) -> str:88 return textwrap.dedent(89 f"""90 Ordering criteria: {obs.ordering_criteria}91 92 Reference tickets:93 {[serialize_ticket(t) for t in obs.reference_tickets]}94 95 Candidate ticket:96 {serialize_ticket(obs.candidate_ticket)}97 98 Existing heuristics:99 { {k: v.model_dump() for k, v in obs.ticket_heuristics.items()} }100 101 Total tickets: {obs.total_tickets}102 Completed iterations: {obs.completed_iterations}103 104 Decide the next action.105 """106 ).strip()107 108 109def get_model_action(client: OpenAI, obs: TicketOrderingObservation) -> Dict[str, Any]:110 user_prompt = build_user_prompt(obs)111 112 try:113 completion = client.chat.completions.create(114 model=MODEL_NAME,115 messages=[116 {"role": "system", "content": SYSTEM_PROMPT},117 {"role": "user", "content": user_prompt},118 ],119 temperature=TEMPERATURE,120 max_tokens=MAX_TOKENS,121 )122 text = (completion.choices[0].message.content or "").strip()123 124 import json125 126 return json.loads(text)127 except Exception:128 return {129 "candidate_priority": backup_rng.uniform(low=0.0, high=1.0),130 "candidate_summary": "issue",131 "next_reference_ids": [],132 "next_candidate_id": backup_rng.choice(list(obs.ticket_heuristics.keys())),133 "end_ordering": False,134 }135 136 137def main() -> None:138 client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)139 with TicketOrderingEnv(base_url=ENV_BASE).sync() as env:140 # async with TicketOrderingEnv(base_url="localhost:8001") as env:141 for task in [GenerationDifficulty.Easy, GenerationDifficulty.Medium, GenerationDifficulty.Hard]:142 rewards: List[float] = []143 steps_taken = 0144 score = 0.0145 success = False146 147 log_start(task=f"ticket_ordering_env_{task.name}", env="ticket_ordering_env", model=MODEL_NAME)148 149 try:150 result = env.reset(difficulty=task.value)151 obs = result.observation152 153 for step in range(1, MAX_STEPS + 1):154 if result.done:155 break156 157 action_dict = get_model_action(client, obs)158 159 action = TicketOrderingAction(160 candidate_priority=float(action_dict.get("candidate_priority", 0.0)),161 candidate_summary=str(action_dict.get("candidate_summary", ""))[:32],162 next_reference_ids=list(action_dict.get("next_reference_ids", [])),163 next_candidate_id=int(action_dict.get("next_candidate_id", 0)),164 end_ordering=bool(action_dict.get("end_ordering", False)),165 )166 167 result = env.step(action)168 obs = result.observation169 170 reward = result.reward or 0.0171 done = result.done172 error = None173 174 rewards.append(reward)175 steps_taken = step176 177 log_step(178 step=step,179 action=str(action_dict),180 reward=reward,181 done=done,182 error=error,183 )184 185 if done:186 break187 188 min_reward = -1.0189 max_reward = 2.0190 rewards_sum = sum(rewards)191 rewards_sum -= min_reward192 rewards_sum /= (max_reward - min_reward)193 score = min(max(rewards_sum, 0.0), 1.0)194 success = score >= SUCCESS_SCORE_THRESHOLD195 196 finally:197 try:198 env.close()199 except Exception:200 pass201 log_end(success=success, steps=steps_taken, score=score, rewards=rewards)202 203 204if __name__ == "__main__":205 main()206 