Spaces:
Running on Zero
Running on Zero
File size: 4,082 Bytes
4a4df15 | 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 | """Versioned, text-only inference protocol. No file paths, model IDs, or arbitrary prompts."""
from __future__ import annotations
import json
import re
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from .schemas import Chunk, Rubric, Concept
from .validation import validate_draft
PROTOCOL = 1
MAX_WIRE_BYTES = 240_000
class Strict(BaseModel):
model_config = ConfigDict(extra='forbid')
class Source(Strict):
id: str = Field(min_length=1,max_length=80,pattern=r'^[A-Za-z0-9_-]+$')
document_id: str = Field(min_length=1,max_length=80)
filename: str = Field(min_length=1,max_length=200)
role: Literal['primary','supporting']
location: str = Field(min_length=1,max_length=180)
heading: str = Field(max_length=300)
text: str = Field(min_length=8,max_length=15000)
ordinal: int = Field(ge=0,le=100000)
class Question(Strict):
topic_label: str = Field(min_length=1,max_length=400)
focus: str = Field(max_length=400,default='')
style: Literal['General description','Definitions and properties','Procedure / operation','Limitations and extensions']
avoid_repeating: list[str] = Field(max_length=4,default_factory=list)
sources: list[Source] = Field(min_length=1,max_length=20)
class Grade(Strict):
question: str = Field(min_length=15,max_length=1300)
answer: str = Field(min_length=1,max_length=18000)
rubric: Rubric
primary_sources: list[Source] = Field(min_length=1,max_length=30)
class Gap(Strict):
concept: Concept
missing_detail: str = Field(max_length=1000)
primary_evidence: list[Source] = Field(min_length=1,max_length=3)
supporting: list[Source] = Field(max_length=3)
class Explain(Strict):
gaps: list[Gap] = Field(min_length=1,max_length=3)
def validate_payload(task: str, payload_json: str) -> dict:
if task not in ('question','grade','explain'):
raise ValueError('Unsupported inference task.')
if not isinstance(payload_json,str) or len(payload_json.encode('utf-8')) > MAX_WIRE_BYTES:
raise ValueError('Inference request is too large.')
try:
data = json.loads(payload_json)
parsed = {'question':Question,'grade':Grade,'explain':Explain}[task].model_validate(data)
except (ValueError,TypeError,ValidationError):
raise ValueError('Invalid inference request structure. Update the desktop package and try again.') from None
groups = []
if task == 'question':
groups = [(parsed.sources,'primary')]
if any(len(x)>1300 for x in parsed.avoid_repeating):
raise ValueError('Question history is too long.')
elif task == 'grade':
groups = [(parsed.primary_sources,'primary')]
if parsed.question != parsed.rubric.question:
raise ValueError('The question must match its frozen rubric.')
validate_draft(parsed.rubric.model_dump(),[Chunk(**s.model_dump()) for s in parsed.primary_sources])
else:
for gap in parsed.gaps:
groups.extend([(gap.primary_evidence,'primary'),(gap.supporting,'supporting')])
for sources,role in groups:
if any(s.role!=role for s in sources) or len({s.id for s in sources}) != len(sources):
raise ValueError('Source IDs must be unique and match their source role.')
if sum(len(s.text) for sources,_ in groups for s in sources)>50000:
raise ValueError('Too much source text. Choose a narrower topic.')
return parsed.model_dump()
def make_response(task: str, result: dict) -> str:
return json.dumps({'protocol':PROTOCOL,'task':task,'result':result},ensure_ascii=False)
def read_response(task: str, raw: str) -> dict:
if not isinstance(raw,str) or len(raw.encode('utf-8'))>MAX_WIRE_BYTES:
raise ValueError('The remote Space returned an invalid response.')
value=json.loads(raw)
if value.get('protocol')!=PROTOCOL or value.get('task')!=task or not isinstance(value.get('result'),dict):
raise ValueError('Remote API version mismatch. Deploy the matching Space package.')
return value['result']
|