"""Decision-model case-routing harness.

Runs the synthetic Salesforce case set (data/salesforce-case-routing-testset.csv)
through one or more decision models and reports accuracy, accuracy on the
ambiguous subset, latency (p50 and p95), and cost at list price.

Published by Prestanda Consulting with the article
https://pcplusa.com/insights/jev-vs-openai-decisions-api-salesforce

Usage:
  pip install requests
  export TYPESAFE_API_KEY=...      # for Jev
  export OPENAI_API_KEY=...        # for GPT-6 Luna via the Responses API
  python3 decision_routing_harness.py --models jev,luna --limit 240

Request shapes follow each vendor's public documentation as of 1 October 2026.
OpenAI's Decisions API is in limited preview with no public reference, so it is
not wired in. Add it when the docs land; the scoring code does not change.
Prices are list prices per million input tokens and are set below; check them
before you trust the cost column. Token counts are estimated at 4 characters
per token where the API does not return usage.
"""
import argparse
import csv
import json
import os
import statistics
import sys
import time

import requests

QUEUES = {
    "billing": "Invoices, charges, refunds, payment failures, tax on invoices",
    "technical_support": "Errors, outages, bugs, broken pages, failed jobs in the product",
    "account_management": "Renewals, contract reviews, churn risk, org changes, named contacts",
    "security": "Suspicious access, breaches, audits, MFA, credentials, phishing",
    "data_integration": "Sync, ERP and middleware flows, bulk loads, duplicates, API limits",
    "sales_upsell": "Pricing requests, more seats, add-ons, trials, proposals",
}
INSTRUCTION = ("Route this Salesforce support case to the queue that should own the "
               "first response.")

PRICE_PER_M_INPUT = {"jev": 0.042, "luna": 0.10}
PRICE_PER_M_OUTPUT = {"jev": 0.0, "luna": 0.50}


def state_for(row):
    return {"subject": row["subject"], "account_tier": row["account_tier"],
            "channel": row["channel"]}


def call_jev(row):
    body = {
        "model": "jev-latest",
        "state": state_for(row),
        "questions": {
            "queue": {"type": "choice", "instructions": INSTRUCTION, "criteria": QUEUES},
        },
    }
    t = time.perf_counter()
    r = requests.post("https://api.typesafe.ai/v1/systemone", timeout=30,
                      headers={"Authorization": f"Bearer {os.environ['TYPESAFE_API_KEY']}"},
                      json=body)
    ms = (time.perf_counter() - t) * 1000
    r.raise_for_status()
    ans = r.json()["answers"]["queue"]
    tokens_in = len(json.dumps(body)) / 4
    return ans["choice"], ans.get("confidence"), ms, tokens_in, 0


def call_luna(row):
    body = {
        "model": "gpt-6-luna",
        "input": [
            {"role": "system", "content": INSTRUCTION + " Queues: " + json.dumps(QUEUES)},
            {"role": "user", "content": json.dumps(state_for(row))},
        ],
        "text": {"format": {"type": "json_schema", "name": "route", "strict": True,
                            "schema": {"type": "object", "additionalProperties": False,
                                       "required": ["queue"],
                                       "properties": {"queue": {"type": "string",
                                                                "enum": list(QUEUES)}}}}},
    }
    t = time.perf_counter()
    r = requests.post("https://api.openai.com/v1/responses", timeout=60,
                      headers={"Authorization": f"Bearer {os.environ['OPENAI_API_KEY']}"},
                      json=body)
    ms = (time.perf_counter() - t) * 1000
    r.raise_for_status()
    data = r.json()
    text = data.get("output_text") or next(
        c["text"] for o in data["output"] if o.get("type") == "message"
        for c in o["content"] if c.get("type") == "output_text")
    usage = data.get("usage", {})
    return (json.loads(text)["queue"], None, ms,
            usage.get("input_tokens", len(json.dumps(body)) / 4),
            usage.get("output_tokens", 8))


CALLERS = {"jev": call_jev, "luna": call_luna}


def pct(values, p):
    values = sorted(values)
    return values[min(len(values) - 1, int(round(p / 100 * (len(values) - 1))))]


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--data", default="salesforce-case-routing-testset.csv")
    ap.add_argument("--models", default="jev,luna")
    ap.add_argument("--limit", type=int, default=240)
    ap.add_argument("--out", default="routing-results.csv")
    args = ap.parse_args()

    rows = list(csv.DictReader(open(args.data)))[: args.limit]
    out = csv.writer(open(args.out, "w", newline=""))
    out.writerow(["model", "case_id", "label", "predicted", "correct", "ambiguous",
                  "confidence", "latency_ms", "tokens_in", "tokens_out"])
    for model in args.models.split(","):
        hits, amb_hits, amb_n, lat, tin, tout, errors = 0, 0, 0, [], 0, 0, 0
        for row in rows:
            try:
                pred, conf, ms, ti, to = CALLERS[model](row)
            except Exception as e:  # record and move on; errors count as misses
                errors += 1
                print(f"{model} {row['case_id']}: {e}", file=sys.stderr)
                pred, conf, ms, ti, to = "ERROR", None, None, 0, 0
            ok = pred == row["label_queue"]
            hits += ok
            if row["ambiguous"] == "yes":
                amb_n += 1
                amb_hits += ok
            if ms is not None:
                lat.append(ms)
            tin += ti
            tout += to
            out.writerow([model, row["case_id"], row["label_queue"], pred, int(ok),
                          row["ambiguous"], conf, round(ms or 0, 1), round(ti), to])
        cost = tin / 1e6 * PRICE_PER_M_INPUT[model] + tout / 1e6 * PRICE_PER_M_OUTPUT[model]
        print(f"\n{model}: accuracy {hits}/{len(rows)} = {hits/len(rows):.1%}"
              f" | ambiguous {amb_hits}/{amb_n}"
              f" | p50 {statistics.median(lat):.0f} ms | p95 {pct(lat, 95):.0f} ms"
              f" | est. cost ${cost:.4f} | errors {errors}")


if __name__ == "__main__":
    main()
