File size: 4,648 Bytes
09d4173
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3c08bbf
09d4173
 
 
 
 
 
 
3c08bbf
09d4173
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b6b0992
 
 
 
 
09d4173
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
144
145
146
147
148
"""Standalone HTTP process; only the backend client talks to an inference server."""

import asyncio
import logging
import secrets
from contextlib import asynccontextmanager

from fastapi import Depends, FastAPI, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse

from .backend import AdapterError, ScoringBackend
from .protocol import SystemOneRequest
from .service import SystemOneService

logger = logging.getLogger(__name__)


def create_app(
    backend: ScoringBackend,
    *,
    api_key: str | None = None,
    max_concurrency: int = 32,
    alias: str = "jev-latest",
    default_temperature: float = 1.0,
    temperature_by_type: dict[str, float] | None = None,
    prompt_wording: str = "served",
) -> FastAPI:
    service = SystemOneService(
        backend,
        max_concurrency=max_concurrency,
        alias=alias,
        default_temperature=default_temperature,
        temperature_by_type=temperature_by_type,
        prompt_wording=prompt_wording,
    )

    @asynccontextmanager
    async def lifespan(app):
        try:
            await backend.start()
            yield
        finally:
            await backend.close()

    app = FastAPI(title="Jev adapter", version="0.1.0", lifespan=lifespan)

    async def authorize(request: Request):
        if api_key and not secrets.compare_digest(
            request.headers.get("authorization", "").encode("utf-8"),
            f"Bearer {api_key}".encode(),
        ):
            raise AdapterError(
                "unauthorized", "A valid bearer token is required.", status=401
            )

    @app.exception_handler(AdapterError)
    async def adapter_error(request, exc):
        return JSONResponse(
            status_code=exc.status,
            content={
                "error": {
                    "code": exc.code,
                    "message": str(exc),
                    "field": exc.field,
                }
            },
        )

    @app.exception_handler(RequestValidationError)
    async def validation_error(request, exc):
        error = exc.errors()[0]
        return JSONResponse(
            status_code=422,
            content={
                "error": {
                    "code": "invalid_request",
                    "message": error["msg"],
                    "field": ".".join(str(p) for p in error["loc"] if p != "body"),
                }
            },
        )

    @app.get("/health")
    async def health():
        # Process liveness, not an inference/GPU readiness probe.
        return {"status": "ok"}

    @app.get("/v1/models", dependencies=[Depends(authorize)])
    async def models():
        return {
            "object": "list",
            "data": [
                {
                    "id": backend.model,
                    "object": "model",
                    "owned_by": "inference-engine",
                    **(
                        {"label_scheme": backend.label_scheme}
                        if getattr(backend, "label_scheme", None)
                        else {}
                    ),
                },
            ],
        }

    @app.post("/v1/systemone", dependencies=[Depends(authorize)])
    async def systemone(body: SystemOneRequest, request: Request):
        work = asyncio.create_task(service.score(body))
        stop_watching = False

        async def disconnected():
            while not stop_watching:
                if await request.is_disconnected():
                    return
                if not stop_watching:
                    await asyncio.sleep(0.05)

        watcher = asyncio.create_task(disconnected())
        try:
            done, _ = await asyncio.wait(
                (work, watcher),
                return_when=asyncio.FIRST_COMPLETED,
            )
            if work in done:
                return await work
            raise AdapterError(
                "client_disconnected", "Client disconnected.", status=499
            )
        except AdapterError:
            raise
        except Exception:
            logger.exception("Decision scoring failed")
            raise AdapterError(
                "scoring_failed",
                "Could not compute a complete decision.",
                status=500,
            ) from None
        finally:
            # A disconnect probe's own cancellation scope can absorb cancel().
            # Stop explicitly as well, including before another polling sleep.
            stop_watching = True
            work.cancel()
            watcher.cancel()
            await asyncio.gather(work, watcher, return_exceptions=True)

    return app