Spaces:
Running on Zero
Running on Zero
Download engine/rpc.py from AngeloUNIMI/document_exam_trainer: direct link, hf CLI and curl.
- Browser
- Download file 4.08 kB
-
https://huggingface.co/spaces/AngeloUNIMI/document_exam_trainer/resolve/main/engine/rpc.py
- Command line
-
hf download hf://spaces/AngeloUNIMI/document_exam_trainer/engine/rpc.py
-
curl -L -o rpc.py https://huggingface.co/spaces/AngeloUNIMI/document_exam_trainer/resolve/main/engine/rpc.py
4.08 kB
| """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'] | |