CoolFace
Apppublic

BHaritha/msme-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
1likes
inference.py254 linesDownload Raw Back to root
1"""2inference.py - MSME Payment Dispute Baseline Agent3 4MANDATORY STDOUT FORMAT (updated spec):5  [START] task=<name> env=msme-dispute model=<model>6  [STEP]  step=<n> action=<json> reward=<0.00> done=<true|false> error=<null|msg>7  [END]   success=<true|false> steps=<n> score=<0.00> rewards=<r1,r2,...>8 9CRITICAL RULES:10  - [END] MUST have score= field11  - All rewards and score must be strictly between 0 and 1 (not 0.00 not 1.00)12  - Use max 0.95, min 0.05 — 0.95 prints as '0.95', 0.999 prints as '1.00' (WRONG)13"""14import os, sys, json, re, requests15 16API_BASE_URL = os.environ.get("API_BASE_URL", "https://api.openai.com/v1")17MODEL_NAME   = os.environ.get("MODEL_NAME",   "gpt-4o-mini")18HF_TOKEN     = os.environ.get("HF_TOKEN",     "")19ENV_URL      = os.environ.get("ENV_URL",       "http://localhost:7860")20ENV_NAME     = "msme-dispute"21 22def _safe(v: float) -> float:23    """Clamp to (0.05, 0.95) — safe for .2f printing."""24    return max(0.05, min(0.95, float(v)))25 26# Fallback actions — used when LLM is unavailable.27# The fallback letter scores ~0.55 from rule-based grader.28FALLBACK = {29    1: {"label": "delayed_payment"},30    2: {31        "claimant":     "Sharma Textiles",32        "opponent":     "Bharat Exports",33        "amount":       80000,34        "due_date":     "31st March 2024",35        "days_overdue": 1436    },37    3: {38        "letter": (39            "April 2024\n\n"40            "To,\nThe Managing Director,\nBharat Exports\n\n"41            "Subject: Legal Demand Notice Under MSMED Act 2006 — "42            "Invoice #1042 — Rs. 80,000\n\n"43            "Dear Sir/Madam,\n\n"44            "This formal legal demand notice is issued by Sharma Textiles against "45            "Bharat Exports for wilful non-payment of Invoice #1042 amounting to "46            "Rs. 80,000 raised on 1st March 2024, with payment due by 31st March 2024. "47            "Despite three written reminders, the outstanding amount remains unpaid.\n\n"48            "Under the MSMED Act 2006 (Micro, Small and Medium Enterprises Development "49            "Act), buyers are legally obligated to clear MSME dues within 45 days of "50            "invoice submission. Your continued default constitutes a clear violation "51            "of the provisions of this Act.\n\n"52            "We hereby demand full payment of Rs. 80,000 within 15 days of receipt "53            "of this notice. In the event of non-payment, compound interest at three "54            "times the RBI bank rate shall be levied on the outstanding amount from "55            "the date of default, as mandated under Section 16 of the MSMED Act 2006.\n\n"56            "We further reserve the right to file a formal complaint before the MSME "57            "Facilitation Council and initiate arbitration under Section 18 of the "58            "MSMED Act 2006 without further notice. All legal costs shall be recovered.\n\n"59            "Kindly treat this as final notice before legal action.\n\n"60            "Yours sincerely,\nShah Sharma\nProprietor, Sharma Textiles"61        )62    }63}64 65# ── Logging ───────────────────────────────────66def log_start(task):67    print(f"[START] task={task} env={ENV_NAME} model={MODEL_NAME}", flush=True)68 69def log_step(step, action, reward, done, error=None):70    a = json.dumps(action, separators=(',',':')).replace('\n',' ')[:200] \71        if isinstance(action, dict) else str(action)[:200]72    r = _safe(reward)73    d = "true" if done else "false"74    e = error if error else "null"75    print(f"[STEP] step={step} action={a} reward={r:.2f} done={d} error={e}", flush=True)76 77def log_end(success, steps, score, rewards):78    """79    [END] format REQUIRES score= field per updated spec.80    score and all rewards must be strictly between 0 and 1.81    """82    sc = _safe(score)83    rs = ",".join(f"{_safe(r):.2f}" for r in rewards)84    s  = "true" if success else "false"85    print(f"[END] success={s} steps={steps} score={sc:.2f} rewards={rs}", flush=True)86 87# ── Env + LLM ─────────────────────────────────88def call_env(ep, payload=None, method="POST"):89    url = f"{ENV_URL}/{ep}"90    r   = requests.get(url, timeout=30) if method == "GET" else \91          requests.post(url, json=payload, timeout=60)92    r.raise_for_status()93    return r.json()94 95def llm(prompt, fallback=""):96    try:97        if not HF_TOKEN or len(HF_TOKEN) < 8:98            raise ValueError("No valid HF_TOKEN")99        from openai import OpenAI100        c = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)101        r = c.chat.completions.create(102            model=MODEL_NAME, max_tokens=1000,103            messages=[{"role": "user", "content": prompt}]104        )105        return r.choices[0].message.content.strip()106    except Exception as e:107        print(f"# LLM error: {e}", file=sys.stderr)108        return fallback109 110# ── Task agents ───────────────────────────────111def agent1(obs):112    raw = llm(113        f"Classify MSME dispute email into ONE label.\n"114        f"Subject: {obs['email']['subject']}\n"115        f"Body: {obs['email']['body']}\n\n"116        f"delayed_payment = full invoice overdue, not paid\n"117        f"partial_payment = only part of invoice paid\n"118        f"payment_denial  = buyer refusing to pay at all\n\n"119        f"Reply with ONLY the label.",120        fallback="delayed_payment"121    ).lower().strip()122    for v in ["delayed_payment", "partial_payment", "payment_denial"]:123        if v in raw: return {"label": v}124    return FALLBACK[1]125 126def agent2(obs):127    raw = llm(128        f"Extract facts from MSME payment notice.\n"129        f"Subject: {obs['email']['subject']}\n"130        f"Body: {obs['email']['body']}\n\n"131        f"Return ONLY valid JSON (no markdown):\n"132        f'{{"claimant":"name","opponent":"name","amount":80000,'133        f'"due_date":"31st March 2024","days_overdue":14}}',134        fallback=json.dumps(FALLBACK[2])135    )136    raw = re.sub(r"```[a-z]*", "", raw).strip().strip("`")137    try:138        r = json.loads(raw)139        if all(k in r for k in ["claimant","opponent","amount","due_date","days_overdue"]):140            return r141    except Exception:142        m = re.search(r"\{.*\}", raw, re.DOTALL)143        if m:144            try: return json.loads(m.group())145            except: pass146    return FALLBACK[2]147 148def agent3(obs):149    ctx = obs["context"]150    try:    amt = f"Rs. {int(ctx.get('amount',0)):,}"151    except: amt = f"Rs. {ctx.get('amount',0)}"152    letter = llm(153        f"Write a formal MSME payment demand letter.\n\n"154        f"Claimant: {ctx.get('claimant')}\n"155        f"Opponent: {ctx.get('opponent')}\n"156        f"Invoice: {ctx.get('invoice_no','N/A')} dated {ctx.get('invoice_date','N/A')}\n"157        f"Due date: {ctx.get('due_date','N/A')}\n"158        f"Amount: {amt}\n"159        f"Days overdue: {ctx.get('days_overdue')}\n"160        f"Dispute: {ctx.get('dispute_type')}\n"161        f"{'' if not obs.get('note') else chr(10) + 'FEEDBACK FROM GRADER TO FIX:' + chr(10) + obs.get('note') + chr(10)}\n"162        f"MUST include ALL of these:\n"163        f"- MSMED Act 2006 (cite explicitly)\n"164        f"- Pay within 15 days\n"165        f"- Compound interest at 3x RBI rate\n"166        f"- Invoice number and amount\n"167        f"- MSME Facilitation Council / legal proceedings\n"168        f"- Assertive legal tone\n"169        f"- Minimum 200 words\n\n"170        f"Write ONLY the letter text.",171        fallback=FALLBACK[3]["letter"]172    )173    if not letter or len(letter.split()) < 50:174        letter = FALLBACK[3]["letter"]175    return {"letter": letter}176 177AGENTS = {1: agent1, 2: agent2, 3: agent3}178NAMES  = {1: "classify_dispute", 2: "extract_facts", 3: "draft_demand_letter"}179 180# ── Run one task episode ──────────────────────181def run_task(task_id: int, seed: int = 42) -> float:182    name = NAMES[task_id]183    log_start(name)184    score = 0.05185    rewards = []186    steps = 0187 188    try:189        resp   = call_env("reset", {"task_id": task_id, "seed": seed})190        obs    = resp["observation"]191        done   = False192        193        while not done and steps < 3:194            steps += 1195            action = AGENTS[task_id](obs)196            result = call_env("step", {"action": action})197            score  = _safe(float(result["reward"]))198            done   = result.get("done", True)199            rewards.append(score)200            log_step(steps, action, score, done)201            202            # Wire multi-turn feedback for task 3203            info = result.get("info", {})204            feedback = info.get("feedback", [])205            message = info.get("message", "")206            if not done and feedback and task_id == 3:207                obs["note"] = f"{message} Missing: " + ", ".join([f["missing"] for f in feedback])208 209    except Exception as e:210        print(f"# Task {task_id} agent failed: {e}", file=sys.stderr)211        try:212            call_env("reset", {"task_id": task_id, "seed": seed})213            result = call_env("step", {"action": FALLBACK[task_id]})214            score  = _safe(float(result["reward"]))215            steps = 1216            rewards = [score]217            log_step(1, FALLBACK[task_id], score, True)218        except Exception as e2:219            score = 0.05220            steps = 1221            rewards = [0.05]222            log_step(1, FALLBACK[task_id], 0.05, True, error=str(e2)[:60])223 224    log_end(True, steps, score, rewards)225    return score226 227# ── Main ──────────────────────────────────────228def main():229    try:230        h = call_env("health", method="GET")231        print(f"# Env: {ENV_URL} | {h.get('status')}", file=sys.stderr)232    except Exception as e:233        print(f"# ERROR: {e}", file=sys.stderr)234        sys.exit(1)235 236    scores = {}237    for t in [1, 2, 3]:238        print(f"\n# Task {t}: {NAMES[t]}", file=sys.stderr)239        scores[t] = run_task(t, seed=42)240 241    print("\n# RESULTS", file=sys.stderr)242    for t, s in scores.items():243        print(f"# Task {t}: {s:.4f}", file=sys.stderr)244    avg = sum(scores.values()) / 3245    print(f"# Average: {avg:.4f}", file=sys.stderr)246 247    os.makedirs("output", exist_ok=True)248    with open("output/inference_results.json", "w") as f:249        json.dump({"task_scores": scores, "model": MODEL_NAME, "env": ENV_URL}, f, indent=2)250    print("# Saved output/inference_results.json", file=sys.stderr)251 252if __name__ == "__main__":253    main()254