File size: 6,020 Bytes
26de23c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""HTTP gateway that speaks OpenRouter's Decisions API in front of an openjev checkpoint served by SGLang (serve_sglang.sh).

    POST /api/alpha/decisions     (also POST /api/v1/systemone: same request/response schema)
    GET  /health                  200 once SGLang answers, 503 before

Env: SGLANG_URL (http://127.0.0.1:30000), API_KEY (Bearer token; unset = no auth, local use only), SERVED_MODEL
(name echoed in the response), PRICE_PER_MTOK (USD per 1M input tokens for usage.cost, default 0), MAX_INFLIGHT (429 above), CLASSIFY_BS.

    uvicorn decisions_server:app --host 0.0.0.0 --port 8000
"""
from __future__ import annotations

import asyncio
import hmac
import json
import os
import random
import string
import time

import httpx
from fastapi import FastAPI, Request
from fastapi.responses import JSONResponse

from decisions_api import ApiError, assemble, build_plan, softmax_ent
from prompt_templates import STYLES, build_prompt_plan, finish_answers
from image_inputs import image_input, image_premise

SGLANG_URL = os.environ.get("SGLANG_URL", "http://127.0.0.1:30000").rstrip("/")
API_KEY = os.environ.get("API_KEY", "")
SERVED_MODEL = os.environ.get("SERVED_MODEL", "openjev/qwen3.5-4b-nli-v5")
PRICE_PER_MTOK = float(os.environ.get("PRICE_PER_MTOK", "0"))
MAX_INFLIGHT = int(os.environ.get("MAX_INFLIGHT", "64"))
CLASSIFY_BS = int(os.environ.get("CLASSIFY_BS", "32"))  # pairs per /classify call; the calls of one request run concurrently
PROMPT_STYLE = os.environ.get("OPENJEV_PROMPT_STYLE", "trained")
if PROMPT_STYLE not in STYLES:
    raise ValueError(f"Unknown OPENJEV_PROMPT_STYLE: {PROMPT_STYLE}")
TEMPLATE = "Premise: {premise}\nHypothesis: {hypothesis}"  # must match train.py / modeling_openjev.py

app = FastAPI(title="openjev decisions", docs_url=None, redoc_url=None, openapi_url=None)
_client: httpx.AsyncClient | None = None
_inflight = 0


def client() -> httpx.AsyncClient:
    global _client
    if _client is None:
        _client = httpx.AsyncClient(timeout=httpx.Timeout(600.0), limits=httpx.Limits(max_connections=256))
    return _client


def err(code: int, message: str) -> JSONResponse:
    return JSONResponse(ApiError(code, message).body(), status_code=code)


async def classify(pairs: list, image=None) -> tuple:
    """-> (P(entailment) per pair, prompt tokens). One request per CLASSIFY_BS pairs, all in flight together."""
    texts = [TEMPLATE.format(premise=image_premise(p.strip()) if image else p.strip(), hypothesis=h.strip())
             for p, h in pairs]

    async def one(chunk):
        payload = {"text": chunk}
        if image:
            payload["image_data"] = [image] * len(chunk)
        r = await client().post(f"{SGLANG_URL}/classify", json=payload)
        r.raise_for_status()
        out = r.json()
        return out if isinstance(out, list) else [out]

    if image:
        # Keep repeated image bytes bounded and vision prefill memory predictable.
        image_bs = max(1, min(4, CLASSIFY_BS, (8 * 1024 * 1024) // len(image)))
        outs = [await one(texts[i:i + image_bs]) for i in range(0, len(texts), image_bs)]
    else:
        outs = await asyncio.gather(*[one(texts[i:i + CLASSIFY_BS]) for i in range(0, len(texts), CLASSIFY_BS)])
    flat = [o for chunk in outs for o in chunk]
    if len(flat) != len(texts):
        raise RuntimeError(f"/classify returned {len(flat)} results for {len(texts)} inputs")
    tokens = sum((o.get("meta_info") or {}).get("prompt_tokens") or 0 for o in flat)
    if not tokens:  # older servers do not report it; ~4 chars per token is close enough for billing display
        tokens = sum(len(t) for t in texts) // 4
    return softmax_ent([o["embedding"] for o in flat]), tokens


def authorized(request: Request) -> bool:
    if not API_KEY:
        return True
    got = request.headers.get("authorization", "")
    return got.lower().startswith("bearer ") and hmac.compare_digest(got[7:].strip(), API_KEY)


@app.get("/health")
async def health():
    try:
        r = await client().get(f"{SGLANG_URL}/health", timeout=3.0)
        if r.status_code == 200:
            return {"status": "ok", "model": SERVED_MODEL}
    except httpx.HTTPError:
        pass
    return JSONResponse({"status": "loading"}, status_code=503)


async def decisions(request: Request):
    global _inflight
    if not authorized(request):
        return err(401, "Missing or invalid Authorization header (expected 'Bearer <key>')")
    try:
        raw = bytearray()
        async for chunk in request.stream():
            raw.extend(chunk)
            if len(raw) > 8 * 1024 * 1024:
                return err(413, "Request exceeds 8 MiB")
        body = json.loads(raw)
    except ValueError:
        return err(400, "request body is not valid JSON")
    try:
        plan = build_prompt_plan(body, PROMPT_STYLE)
        image = image_input(body)
        if image:
            for premise, _ in plan.pairs:
                image_premise(premise)
    except ApiError as e:
        return err(e.code, e.message)
    if _inflight >= MAX_INFLIGHT:
        return err(429, "Too many concurrent requests")
    _inflight += 1
    try:
        ent, tokens = await classify(plan.pairs, image=image)
        answers = finish_answers(plan, ent)
    except (httpx.HTTPError, RuntimeError, KeyError, ValueError) as e:
        return err(502, f"upstream model server failed: {type(e).__name__}: {str(e)[:200]}")
    finally:
        _inflight -= 1
    rid = "gen-dec-%d-%s" % (time.time(), "".join(random.choices(string.ascii_letters + string.digits, k=20)))
    return {"id": rid, "model": SERVED_MODEL, "provider": "openjev", "answers": answers,
            "probability_method": "normalized_entailment_v1",
            "usage": {"input_tokens": int(tokens), "output_tokens": 0, "cost": tokens * PRICE_PER_MTOK / 1e6}}


app.add_api_route("/api/alpha/decisions", decisions, methods=["POST"])
app.add_api_route("/api/v1/systemone", decisions, methods=["POST"])
app.add_api_route("/v1/systemone", decisions, methods=["POST"])