diff --git a/README.md b/README.md index fafc79b..f1b6903 100644 --- a/README.md +++ b/README.md @@ -9,10 +9,102 @@ pinned: false # Contract Negotiation Environment -An AI-powered OpenEnv environment for contract negotiation using hybrid rule-based and LLM-driven decision making. +An AI-powered OpenEnv environment for evaluating contract-negotiation agents +through hybrid rule-based and LLM-driven decision making. ## Overview -This project simulates real-world contract negotiation with structured tasks, grading, and agent evaluation. + +This project simulates real-world contract negotiation scenarios where an AI +agent must: + +1. **Analyse** contract clauses to identify legal risks (unlimited liability, + hidden traps, one-sided IP terms, etc.). +2. **Decide** on the best negotiation action: flag the risk, edit the clause, + propose a counter-offer, reject, or accept. +3. **Generate** safer clause rewrites that protect the customer while keeping + commercially reasonable terms. + +Agents are scored on three dimensions: +- **Correctness** — how well the agent identifies risky language. +- **Improvement** — how much the proposed edits reduce risk. +- **Risk alignment** — whether the chosen action matches the actual risk level. + +## Tasks + +| ID | Difficulty | Clause Type | Industry | +|----|-----------|-------------|----------| +| `easy_unlimited_liability` | EASY | Liability | SaaS B2B | +| `medium_auto_renewal` | MEDIUM | Term/Renewal | SaaS B2B | +| `hard_conflicting_obligations` | HARD | Performance/Changes | Professional Services | +| `easy_compliance_agreement` | EASY+ | Compliance | SaaS B2B | +| `hard_intellectual_property` | HARD+ | IP Ownership | Professional Services | + +Each task has a **dedicated grader** with difficulty-specific scoring adjustments +(e.g., harder tasks penalise unresolved hidden traps). ## Structure -- `contract_env/` → core environment, server, inference \ No newline at end of file + +``` +contract_env/ +├── env/ +│ ├── environment.py # ContractEnv — the main OpenEnv environment +│ ├── graders.py # Task-specific grading functions +│ ├── models.py # Pydantic models (Action, Reward, Observation) +│ └── tasks.py # Task definitions and metadata +├── server/ +│ └── app.py # FastAPI server exposing /reset, /step, /state, /tasks +├── tests/ # Unit tests for API, graders, and environment +└── scripts/ # Helper scripts for local/Docker runs +inference.py # LLM-driven inference agent +openenv.yaml # OpenEnv manifest +Dockerfile # Production container definition +``` + +## Quick Start + +### Local development + +```bash +pip install -e ".[dev]" +python -m pytest contract_env/tests/ -v +``` + +### Run the server + +```bash +uvicorn contract_env.server.app:app --host 0.0.0.0 --port 7860 +``` + +### Run inference + +```bash +export HF_TOKEN="your-huggingface-token" +python inference.py --episodes 5 +python inference.py --benchmark # one episode per task +``` + +### Docker + +```bash +docker build -t contract-negotiation-env . +docker run -p 7860:7860 contract-negotiation-env +``` + +## API Endpoints + +| Method | Path | Description | +|--------|------|-------------| +| `GET` | `/health` | Health check | +| `GET` | `/tasks` | List all tasks with metadata | +| `GET` | `/state` | Current environment state | +| `POST` | `/reset` | Reset and get first observation | +| `POST` | `/step` | Submit an action, receive reward | + +## Environment Variables + +| Variable | Required | Default | Description | +|----------|----------|---------|-------------| +| `API_BASE_URL` | No | `https://router.huggingface.co/v1` | LLM API endpoint | +| `MODEL_NAME` | No | `Qwen/Qwen2.5-72B-Instruct` | Model identifier | +| `HF_TOKEN` | Yes | — | HuggingFace API token | +| `PORT` | No | `7860` | Server port | \ No newline at end of file diff --git a/contract_env/env/environment.py b/contract_env/env/environment.py index d06b404..d252268 100644 --- a/contract_env/env/environment.py +++ b/contract_env/env/environment.py @@ -102,11 +102,17 @@ def step(self, action: Action) -> Tuple[Observation, float, bool, dict[str, Any] contract_before = self.state_data["contract_text"] proposed = build_proposed_contract_for_step(contract_before, action) - reward_obj, grade_info = evaluate_action( - self.current_task, contract_before, action, proposed - ) + # Use the task-specific grader when available + task = self.current_task + if task.has_grader(): + reward_obj = task.grade(contract_before, action, proposed) + # Collect grade info from evaluate_action for transparency + _, grade_info = evaluate_action(task, contract_before, action, proposed) + else: + reward_obj, grade_info = evaluate_action( + task, contract_before, action, proposed + ) - # ✅ FIX: convert Reward → float reward = float(reward_obj.score) info.update(grade_info) diff --git a/contract_env/env/graders.py b/contract_env/env/graders.py index f2368bc..d97d764 100644 --- a/contract_env/env/graders.py +++ b/contract_env/env/graders.py @@ -182,25 +182,68 @@ def grade_action( return reward -# Specific graders for each task +# ============ TASK-SPECIFIC GRADERS ============ +# Each grader applies difficulty-specific adjustments on top of the base evaluation. + +# -- Grading multipliers (named constants for clarity) -- +_EASY_SAFE_EDIT_BONUS = 1.08 # +8 % for well-matched safe edits +_MEDIUM_PREMATURE_ACCEPT_PENALTY = 0.65 # −35 % for accepting risky terms +_HARD_UNRESOLVED_TRAP_PENALTY = 0.5 # −50 % when hidden traps remain +_EASY_PLUS_NOTIFICATION_BONUS = 1.06 # +6 % for breach-notification language +_HARD_PLUS_TRAP_PENALTY = 0.55 # −45 % for unresolved IP traps +_HARD_PLUS_OWNERSHIP_BONUS = 1.07 # +7 % for explicit customer-ownership + + def grade_easy(task: NegotiationTask, contract_before: str, action: Action, proposed_contract_text: str) -> Reward: - return grade_action(task, contract_before, action, proposed_contract_text) + """Grade easy tasks with a bias toward accepting safe-looking clauses quickly.""" + reward, _ = evaluate_action(task, contract_before, action, proposed_contract_text) + if action.action_type in ("EDIT_CLAUSE", "PROPOSE_COUNTER"): + safe = _safe_overlap( + (action.content or "").strip(), task.safe_keywords, task.expected_safe_edit + ) + if safe > 0.5: + reward.score = max(0.001, min(0.999, reward.score * _EASY_SAFE_EDIT_BONUS)) + return reward def grade_medium(task: NegotiationTask, contract_before: str, action: Action, proposed_contract_text: str) -> Reward: - return grade_action(task, contract_before, action, proposed_contract_text) + """Grade medium tasks, penalising premature acceptance of risky auto-renewal terms.""" + reward, _ = evaluate_action(task, contract_before, action, proposed_contract_text) + if action.action_type == "ACCEPT": + risk = _weighted_risk_hits(proposed_contract_text, task.risk_keywords) + if risk >= 0.3: + reward.score = max(0.001, min(0.999, reward.score * _MEDIUM_PREMATURE_ACCEPT_PENALTY)) + return reward def grade_hard(task: NegotiationTask, contract_before: str, action: Action, proposed_contract_text: str) -> Reward: - return grade_action(task, contract_before, action, proposed_contract_text) + """Grade hard tasks with trap-resolution checking and heavier penalty for missed traps.""" + reward, _ = evaluate_action(task, contract_before, action, proposed_contract_text) + if action.action_type in ("ACCEPT", "EDIT_CLAUSE", "PROPOSE_COUNTER"): + if trap_unresolved(task, proposed_contract_text): + reward.score = max(0.001, min(0.999, reward.score * _HARD_UNRESOLVED_TRAP_PENALTY)) + return reward def grade_easy_plus(task: NegotiationTask, contract_before: str, action: Action, proposed_contract_text: str) -> Reward: - return grade_action(task, contract_before, action, proposed_contract_text) + """Grade easy-plus compliance tasks, rewarding mention of notification obligations.""" + reward, _ = evaluate_action(task, contract_before, action, proposed_contract_text) + content = (action.content or "").strip().lower() + if action.action_type in ("EDIT_CLAUSE", "PROPOSE_COUNTER"): + if any(kw in content for kw in ("notify", "notification", "promptly inform")): + reward.score = max(0.001, min(0.999, reward.score * _EASY_PLUS_NOTIFICATION_BONUS)) + return reward def grade_hard_plus(task: NegotiationTask, contract_before: str, action: Action, proposed_contract_text: str) -> Reward: - return grade_action(task, contract_before, action, proposed_contract_text) + """Grade hard-plus IP tasks with trap-resolution + ownership-clarity checks.""" + reward, _ = evaluate_action(task, contract_before, action, proposed_contract_text) + content = (action.content or "").strip().lower() + if trap_unresolved(task, proposed_contract_text): + reward.score = max(0.001, min(0.999, reward.score * _HARD_PLUS_TRAP_PENALTY)) + if any(kw in content for kw in ("customer owns", "customer-owned", "owned by customer")): + reward.score = max(0.001, min(0.999, reward.score * _HARD_PLUS_OWNERSHIP_BONUS)) + return reward # ============ GRADER REGISTRY ============ diff --git a/contract_env/server/app.py b/contract_env/server/app.py index fd7d709..d791e6f 100644 --- a/contract_env/server/app.py +++ b/contract_env/server/app.py @@ -11,13 +11,19 @@ from pydantic import ValidationError from contract_env.env.environment import ContractEnv +from contract_env.env.graders import TASK_GRADERS, NUM_GRADED_TASKS from contract_env.env.models import Action, StepRequest +from contract_env.env.tasks import TASKS _env = ContractEnv() app = FastAPI( title="Contract Negotiation OpenEnv", - version="1.1.0", + description=( + "AI-driven environment for evaluating contract-negotiation agents. " + "Agents analyse clauses, identify risks, and propose safer alternatives." + ), + version="1.2.0", ) app.add_middleware( @@ -29,13 +35,13 @@ ) -# ---------------- ROOT ---------------- +# ── ROOT ──────────────────────────────────────────────────────────────── @app.get("/") def root(): return {"status": "ok", "service": "contract-negotiation-env"} -# ---------------- ERROR HANDLERS ---------------- +# ── ERROR HANDLERS ────────────────────────────────────────────────────── @app.exception_handler(HTTPException) async def http_exception_handler(request: Request, exc: HTTPException): return JSONResponse( @@ -60,43 +66,59 @@ async def global_exception_handler(request: Request, exc: Exception): ) -# ---------------- HEALTH ---------------- +# ── HEALTH ────────────────────────────────────────────────────────────── @app.get("/health") def health(): return {"status": "ok"} -# ---------------- STATE ---------------- +# ── TASKS LISTING ─────────────────────────────────────────────────────── +@app.get("/tasks") +def list_tasks(): + """Return metadata for every registered task.""" + return { + "total": len(TASKS), + "graded": NUM_GRADED_TASKS, + "tasks": [ + { + "id": t.id, + "name": t.name, + "clause_type": t.clause_type, + "risk_level": t.risk_level, + "industry_context": t.industry_context, + "has_grader": t.id in TASK_GRADERS, + } + for t in TASKS + ], + } + + +# ── STATE ─────────────────────────────────────────────────────────────── @app.get("/state") def get_state(): return _env.state() -# ---------------- RESET (FIXED) ---------------- +# ── RESET ─────────────────────────────────────────────────────────────── @app.post("/reset") def reset(): try: obs = _env.reset() - - return { - "observation": obs.model_dump() - } - + return {"observation": obs.model_dump()} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) -# ---------------- STEP (FIXED) ---------------- +# ── STEP ──────────────────────────────────────────────────────────────── @app.post("/step") def step(req: StepRequest): try: action = Action(action_type=req.action_type, content=req.content) - obs, reward, done, info = _env.step(action) return { "observation": obs.model_dump(), - "reward": {"score": reward}, # CRITICAL FIX + "reward": {"score": reward}, "done": done, "info": info, } @@ -105,7 +127,7 @@ def step(req: StepRequest): raise HTTPException(status_code=422, detail=e.errors()) -# ---------------- MAIN ---------------- +# ── MAIN ──────────────────────────────────────────────────────────────── def main(): import uvicorn diff --git a/contract_env/tests/test_api.py b/contract_env/tests/test_api.py index 2a49f57..222ed0f 100644 --- a/contract_env/tests/test_api.py +++ b/contract_env/tests/test_api.py @@ -43,6 +43,18 @@ def test_state_endpoint(self) -> None: self.assertEqual(r.status_code, 200) self.assertIn("current_step", r.json()) + def test_tasks_endpoint(self) -> None: + r = self.client.get("/tasks") + self.assertEqual(r.status_code, 200) + data = r.json() + self.assertGreaterEqual(data["total"], 5) + self.assertGreaterEqual(data["graded"], 3) + self.assertEqual(len(data["tasks"]), data["total"]) + for t in data["tasks"]: + self.assertIn("id", t) + self.assertIn("clause_type", t) + self.assertIn("has_grader", t) + if __name__ == "__main__": unittest.main() diff --git a/contract_env/tests/test_graders.py b/contract_env/tests/test_graders.py index 1a4a2a0..c4a1464 100644 --- a/contract_env/tests/test_graders.py +++ b/contract_env/tests/test_graders.py @@ -7,6 +7,11 @@ effective_risk_high, evaluate_action, grade_action, + grade_easy, + grade_medium, + grade_hard, + grade_easy_plus, + grade_hard_plus, token_overlap_ratio, ) from contract_env.env.models import Action @@ -59,6 +64,50 @@ def test_effective_high_hard_trap(self) -> None: effective_risk_high(task, task.expected_safe_edit), ) + # ── Differentiated grader tests ───────────────────────────────────── + def test_grade_easy_rewards_safe_edit(self) -> None: + task = next(t for t in TASKS if t.name == "EASY") + action = Action(action_type="EDIT_CLAUSE", content=task.expected_safe_edit) + r = grade_easy(task, task.contract_text, action, task.expected_safe_edit) + self.assertGreater(r.score, 0.0) + self.assertLess(r.score, 1.0) + + def test_grade_medium_penalises_premature_accept(self) -> None: + task = next(t for t in TASKS if t.name == "MEDIUM") + r = grade_medium(task, task.contract_text, Action(action_type="ACCEPT"), task.contract_text) + self.assertLessEqual(r.score, 0.01) + + def test_grade_hard_penalises_unresolved_trap(self) -> None: + task = next(t for t in TASKS if t.name == "HARD") + # Accepting original text with traps should score low + action = Action(action_type="EDIT_CLAUSE", content=task.contract_text) + r = grade_hard(task, task.contract_text, action, task.contract_text) + r_safe = grade_hard(task, task.contract_text, + Action(action_type="EDIT_CLAUSE", content=task.expected_safe_edit), + task.expected_safe_edit) + self.assertGreater(r_safe.score, r.score) + + def test_grade_easy_plus_bounds(self) -> None: + task = next(t for t in TASKS if t.name == "EASY_PLUS") + a = Action(action_type="FLAG_RISK", content="note") + r = grade_easy_plus(task, task.contract_text, a, task.contract_text) + self.assertGreater(r.score, 0.0) + self.assertLess(r.score, 1.0) + + def test_grade_hard_plus_penalises_unresolved_trap(self) -> None: + task = next(t for t in TASKS if t.name == "HARD_PLUS") + # Edit that keeps trap markers should score lower than safe edit + action = Action(action_type="EDIT_CLAUSE", content=task.contract_text) + r_trap = grade_hard_plus(task, task.contract_text, action, task.contract_text) + r_safe = grade_hard_plus(task, task.contract_text, + Action(action_type="EDIT_CLAUSE", content=task.expected_safe_edit), + task.expected_safe_edit) + self.assertGreater(r_safe.score, r_trap.score) + + def test_all_tasks_have_graders(self) -> None: + for task in TASKS: + self.assertTrue(task.has_grader(), f"Task {task.id} missing grader") + if __name__ == "__main__": unittest.main() diff --git a/inference.py b/inference.py index 479c837..34416af 100644 --- a/inference.py +++ b/inference.py @@ -1,11 +1,23 @@ +""" +Inference Script — Contract Negotiation Environment +===================================================== +LLM-driven agent that analyses contract clauses, identifies legal risks, +and proposes safer alternatives through multi-turn negotiation. + +MANDATORY environment variables + API_BASE_URL The API endpoint for the LLM. + MODEL_NAME The model identifier to use for inference. + HF_TOKEN Your Hugging Face / API key. +""" from __future__ import annotations import argparse import json +import logging import os import random +import re import sys -import urllib.request import warnings from typing import Any, Optional @@ -23,32 +35,96 @@ warnings.filterwarnings("ignore") -# ---------------- ENV CONFIG ---------------- -API_BASE_URL = os.environ.get("API_BASE_URL", "https://router.huggingface.co/v1") +log = logging.getLogger(__name__) + +# ── ENV CONFIG ────────────────────────────────────────────────────────── +API_BASE_URL = os.environ.get( + "API_BASE_URL", "https://router.huggingface.co/v1" +) MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct") HF_TOKEN = os.getenv("HF_TOKEN") or os.getenv("API_KEY") +# ── LLM CLIENT (lazy singleton) ──────────────────────────────────────── +_client: Optional[OpenAI] = None + + +def _get_client() -> OpenAI: + global _client + if _client is None: + _client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN) + return _client + + +# ── SYSTEM PROMPT ─────────────────────────────────────────────────────── +SYSTEM_PROMPT = """\ +You are an expert contract-negotiation AI assistant working for the Customer. + +Your goals — in priority order: +1. Identify every legal risk, hidden trap, or one-sided obligation. +2. Propose concrete, balanced rewrites that cap liability, ensure mutual + obligations, add reasonable notice periods, and clarify IP ownership. +3. Only ACCEPT a clause once all material risks are resolved. + +When analysing a clause you MUST return **valid JSON** with the schema: +{ + "risk_assessment": "", + "risk_level": "HIGH" | "MODERATE" | "LOW", + "recommended_action": "FLAG_RISK" | "EDIT_CLAUSE" | "PROPOSE_COUNTER" | "REJECT" | "ACCEPT", + "rewritten_clause": "" +} + +Rules: +- Never accept unlimited liability, one-day notice periods, or clauses that + assign all IP to the supplier when customer provides specifications. +- Prefer EDIT_CLAUSE when you can rewrite the clause directly. +- Use PROPOSE_COUNTER when a full counter-offer is warranted. +- Use REJECT only for egregiously one-sided terms that cannot be edited. +- Use FLAG_RISK as the first move for HIGH-risk clauses before editing. +- Return ONLY the JSON object, no markdown fences, no commentary. +""" -# ---------------- HTTP ---------------- -def _http_post(path: str, payload: Optional[dict] = None) -> dict[str, Any]: +# ── LLM HELPERS ───────────────────────────────────────────────────────── +_MAX_RETRIES = 2 + + +def _llm_chat(messages: list[dict], temperature: float = 0.15, + max_tokens: int = 512) -> str: + """Call the LLM with retry logic. Returns the raw text response.""" + client = _get_client() + for attempt in range(_MAX_RETRIES + 1): + try: + resp = client.chat.completions.create( + model=MODEL_NAME, + messages=messages, + temperature=temperature, + max_tokens=max_tokens, + ) + return (resp.choices[0].message.content or "").strip() + except Exception as exc: + if attempt == _MAX_RETRIES: + raise + log.warning("LLM call attempt %d failed: %s", attempt + 1, exc) + return "" + + +def _parse_llm_json(text: str) -> Optional[dict]: + """Best-effort extraction of a JSON object from LLM output.""" + # Strip markdown code fences if present + cleaned = re.sub(r"```(?:json)?", "", text).strip().rstrip("`") try: - url = f"{API_BASE_URL}{path}" - data = None - headers = {"Content-Type": "application/json"} - - if payload is not None: - data = json.dumps(payload).encode("utf-8") - - req = urllib.request.Request(url, data=data, headers=headers, method="POST") - - with urllib.request.urlopen(req, timeout=30) as resp: - return json.loads(resp.read().decode("utf-8")) - - except Exception as e: - raise RuntimeError(f"HTTP request failed: {str(e)}") + return json.loads(cleaned) + except json.JSONDecodeError: + # Try to find JSON object in the text + match = re.search(r"\{[\s\S]*\}", cleaned) + if match: + try: + return json.loads(match.group()) + except json.JSONDecodeError: + pass + return None -# ---------------- SCORING ---------------- +# ── RISK ANALYSIS (rule-based fallback) ───────────────────────────────── def _risk_score(task: NegotiationTask, contract_text: str) -> float: hits = keyword_match_score(contract_text, task.risk_keywords) rs = min(1.0, hits * task.clause_type_weight / 1.15) @@ -57,174 +133,205 @@ def _risk_score(task: NegotiationTask, contract_text: str) -> float: return round(rs, 6) -def _confidence_and_intent(task: NegotiationTask, contract_text: str): +def _rule_based_intent(task: NegotiationTask, contract_text: str) -> str: rs = _risk_score(task, contract_text) - confidence = max(0.0, min(1.0, 1.0 - rs)) - if effective_risk_high(task, contract_text) or rs >= 0.6: - return confidence, "HIGH" + return "HIGH" if task.risk_level.upper() == "MODERATE" and rs >= 0.35: - return confidence, "MODERATE" - return confidence, "LOW" - - -# ---------------- ACTION ---------------- -def _content_for(task: NegotiationTask, action_type: str): - if action_type in ("EDIT_CLAUSE", "PROPOSE_COUNTER"): - return task.expected_safe_edit - return None - - -def _action_for(task: NegotiationTask, action_type: str): - return Action(action_type=action_type, content=_content_for(task, action_type)) - - -# ---------------- 🔥 FORCE LLM CALL ---------------- -def _force_llm_call(contract_text: str): - try: - client = OpenAI( - base_url=API_BASE_URL, - api_key=HF_TOKEN, - ) - - resp = client.chat.completions.create( - model=MODEL_NAME, - messages=[ - {"role": "user", "content": f"Analyze this contract:\n{contract_text}"} - ], - max_tokens=50, - ) + return "MODERATE" + return "LOW" + + +# ── LLM-DRIVEN STRATEGY ──────────────────────────────────────────────── +def _build_analysis_prompt(task: NegotiationTask, state_data: dict, + step: int, history_summary: str) -> list[dict]: + """Build the chat messages for the LLM analysis call.""" + user_msg = ( + f"Contract clause (type: {task.clause_type}, " + f"industry: {task.industry_context}):\n" + f'"""\n{state_data["contract_text"]}\n"""\n\n' + ) + if history_summary: + user_msg += f"Negotiation history so far:\n{history_summary}\n\n" + user_msg += ( + f"This is negotiation step {step + 1}. " + "Analyse the clause and return your JSON recommendation." + ) + return [ + {"role": "system", "content": SYSTEM_PROMPT}, + {"role": "user", "content": user_msg}, + ] + + +def _build_rewrite_prompt(task: NegotiationTask, + contract_text: str, + risk_assessment: str) -> list[dict]: + """Build a focused rewrite prompt when the analysis step doesn't return + a usable rewritten_clause.""" + user_msg = ( + f"You previously identified these risks in this {task.clause_type} " + f"clause:\n{risk_assessment}\n\n" + f"Original clause:\n{contract_text}\n\n" + "Rewrite the clause to eliminate all identified risks while keeping " + "reasonable commercial terms. Return ONLY the rewritten clause text, " + "nothing else." + ) + return [ + {"role": "system", "content": SYSTEM_PROMPT}, + {"role": "user", "content": user_msg}, + ] - _ = resp.choices[0].message.content - except Exception: - pass +_VALID_ACTIONS = {"FLAG_RISK", "EDIT_CLAUSE", "ACCEPT", "REJECT", + "PROPOSE_COUNTER"} -# ---------------- LLM IMPROVEMENT ---------------- -def _maybe_llm_improve(task, contract_text, action): - if action.action_type not in ("EDIT_CLAUSE", "PROPOSE_COUNTER"): - return action +def _choose(task: NegotiationTask, state_data: dict, step: int, + prev_rewards: list[float]) -> Action: + """Use the LLM to decide the next action, falling back to rules on error.""" + contract_text = state_data["contract_text"] + history = state_data.get("negotiation_history", []) + history_summary = "\n".join(history[-6:]) if history else "" + # ── 1. Ask the LLM for a structured analysis ─────────────────────── try: - client = OpenAI( - base_url=API_BASE_URL, - api_key=HF_TOKEN, - ) - - prompt = f""" -You are an AI contract negotiation assistant. - -Given this clause: -{action.content or contract_text} - -Rewrite it to make it safer and reduce legal risk. - -Return ONLY the improved clause text. -""" - - resp = client.chat.completions.create( - model=MODEL_NAME, - messages=[{"role": "user", "content": prompt}], - temperature=0.2, - max_tokens=256, - ) - - text = (resp.choices[0].message.content or "").strip() - - if text: - return Action(action_type=action.action_type, content=text) - - except Exception: - pass - - return action - - -# ---------------- STRATEGY ---------------- -def _sequence(intent: str): - if intent == "HIGH": - return ["FLAG_RISK", "REJECT", "PROPOSE_COUNTER", "EDIT_CLAUSE", "ACCEPT"] - if intent == "MODERATE": - return ["FLAG_RISK", "EDIT_CLAUSE", "PROPOSE_COUNTER", "ACCEPT"] - return ["ACCEPT", "EDIT_CLAUSE", "PROPOSE_COUNTER"] - - -def _choose(task, state_data, step): - confidence, intent = _confidence_and_intent(task, state_data["contract_text"]) - seq = _sequence(intent) - - action_type = seq[min(step, len(seq) - 1)] - action = _action_for(task, action_type) - - # 🔥 GUARANTEE at least ONE LLM call - if step == 0: - _force_llm_call(state_data["contract_text"]) - - return _maybe_llm_improve(task, state_data["contract_text"], action) - - -# ---------------- LOGGING ---------------- -def _log_step(step, action, reward, done, err): + messages = _build_analysis_prompt(task, state_data, step, + history_summary) + raw = _llm_chat(messages) + parsed = _parse_llm_json(raw) + except Exception as exc: + log.warning("LLM analysis call failed: %s", exc) + parsed = None + + # ── 2. Extract action + content from the LLM response ────────────── + action_type: Optional[str] = None + content: Optional[str] = None + risk_assessment: str = "" + + if parsed: + rec = (parsed.get("recommended_action") or "").upper().strip() + if rec in _VALID_ACTIONS: + action_type = rec + content = parsed.get("rewritten_clause") or None + risk_assessment = parsed.get("risk_assessment", "") + + # ── 3. Rule-based fallback if LLM didn't return valid action ─────── + if action_type is None: + intent = _rule_based_intent(task, contract_text) + if intent == "HIGH": + seq = ["FLAG_RISK", "EDIT_CLAUSE", "PROPOSE_COUNTER", "REJECT", + "ACCEPT"] + elif intent == "MODERATE": + seq = ["FLAG_RISK", "EDIT_CLAUSE", "PROPOSE_COUNTER", "ACCEPT"] + else: + seq = ["EDIT_CLAUSE", "PROPOSE_COUNTER", "ACCEPT"] + action_type = seq[min(step, len(seq) - 1)] + + # ── 4. Adaptive adjustment based on previous reward feedback ─────── + if prev_rewards and prev_rewards[-1] < 0.2 and step > 0: + # Previous action scored poorly — try editing instead of repeating + if action_type in ("FLAG_RISK", "REJECT"): + action_type = "EDIT_CLAUSE" + + # ── 5. Generate content for EDIT / PROPOSE if missing ────────────── + if action_type in ("EDIT_CLAUSE", "PROPOSE_COUNTER") and not content: + try: + msgs = _build_rewrite_prompt(task, contract_text, + risk_assessment or "High legal risk") + content = _llm_chat(msgs, max_tokens=384) + # Strip any quotes the model might wrap around + if content.startswith('"') and content.endswith('"'): + content = content[1:-1] + except Exception as exc: + log.warning("LLM rewrite call failed: %s", exc) + content = None + + # ── 6. Ensure content actions always have content ────────────────── + if action_type in ("EDIT_CLAUSE", "PROPOSE_COUNTER") and not content: + content = task.expected_safe_edit # safe fallback + + return Action(action_type=action_type, content=content) + + +# ── LOGGING ───────────────────────────────────────────────────────────── +def _log_step(step: int, action: Action, reward: float, done: bool, + err: Optional[str]) -> None: err_token = "null" if not err else err print( - f"[STEP] step={step} action={action.action_type} reward={reward:.2f} done={str(done).lower()} error={err_token}", + f"[STEP] step={step} action={action.action_type} " + f"reward={reward:.2f} done={str(done).lower()} error={err_token}", flush=True, ) -# ---------------- EXECUTION ---------------- -def run_episode(): +# ── EPISODE EXECUTION ────────────────────────────────────────────────── +def run_episode() -> None: env = ContractEnv() obs = env.reset().model_dump(mode="json") task = next(t for t in TASKS if t.clause_type == obs["clause_type"]) - state_data = { + state_data: dict[str, Any] = { "contract_text": obs["contract_text"], "negotiation_history": list(obs.get("negotiation_history", [])), } - print(f"[START] task={task.name} env=ContractNegotiationEnv model={MODEL_NAME}", flush=True) + print( + f"[START] task={task.name} env=ContractNegotiationEnv " + f"model={MODEL_NAME}", + flush=True, + ) - rewards = [] + rewards: list[float] = [] done = False step = 0 try: while not done and step < 10: - action = _choose(task, state_data, step) + action = _choose(task, state_data, step, rewards) - obs, reward, done, info = env.step(action) + obs_obj, reward, done, info = env.step(action) score = float(reward) rewards.append(score) - state_data["contract_text"] = obs.contract_text + state_data["contract_text"] = obs_obj.contract_text + state_data["negotiation_history"] = list( + obs_obj.negotiation_history + ) _log_step(step, action, score, done, info.get("error")) - step += 1 final_score = sum(rewards) / max(len(rewards), 1) - rewards_str = ",".join(f"{r:.2f}" for r in rewards) print( - f"[END] success={str(final_score >= 0.5).lower()} steps={step} score={final_score:.2f} rewards={rewards_str}", flush=True + f"[END] success={str(final_score >= 0.5).lower()} " + f"steps={step} score={final_score:.2f} rewards={rewards_str}", + flush=True, ) except Exception as e: - print(f'[STEP] step=0 action=NONE reward=0.00 done=true error={str(e)}', flush=True) + print( + f"[STEP] step=0 action=NONE reward=0.00 " + f"done=true error={str(e)}", + flush=True, + ) print("[END] success=false steps=0 score=0.00 rewards=", flush=True) return -def main(): +# ── MAIN ──────────────────────────────────────────────────────────────── +def main() -> None: load_dotenv() random.seed(42) - parser = argparse.ArgumentParser() - parser.add_argument("--episodes", type=int, default=5) - parser.add_argument("--benchmark", action="store_true") + parser = argparse.ArgumentParser( + description="Run contract-negotiation inference episodes", + ) + parser.add_argument("--episodes", type=int, default=5, + help="Number of episodes to run") + parser.add_argument("--benchmark", action="store_true", + help="Run one episode per task") args = parser.parse_args() episodes_to_run = len(TASKS) if args.benchmark else args.episodes