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']