botp
/

Solomon / src /solomon /service_checked.py
orz99's picture ArcherHume's picture
Duplicate from DoccyHealth/Solomon
1d2de8a
Raw History Blame Contribute Delete
7.6 kB
"""Backend-neutral contract-v3 service; restart replays immutable inputs, not tensors.
Engine protocol: identity (mapping or method) -> JSON mapping with fingerprint; prefill(parts) -> state;
ask(state, block, n, execution='cached'|'full') -> letter_logits + branch_tokens.
Engine.ask must fork its prefix cache for every branch. One service serializes all
engine calls. A service instance exclusively owns its engine.
"""
import copy
import hashlib
import json
import re
import shutil
import uuid
from pathlib import Path
import numpy as np
from solomon.service_states import Service as ReadoutService, handler
def _json(value):
return json.dumps(value, sort_keys=True, separators=(',', ':'), allow_nan=False)
def _identity(engine):
value = engine.identity() if callable(engine.identity) else engine.identity
if not isinstance(value, dict) or not isinstance(value.get('fingerprint'), str) or not value['fingerprint']:
raise ValueError('engine identity requires a nonempty fingerprint')
return json.loads(_json(value))
def _parts(document):
if isinstance(document, str):
document = [{'text': document}]
if not isinstance(document, list) or not document:
raise ValueError('document must be text or a nonempty parts list')
result = []
for part in document:
if not isinstance(part, dict) or set(part) not in ({'text'}, {'image'}):
raise ValueError('each part must contain exactly text or image')
name, value = next(iter(part.items()))
if not isinstance(value, str) or not value.strip():
raise ValueError('document parts must contain nonempty strings')
result.append({name: value})
return result
def _digest(parts):
# JSON framing preserves boundaries; two text parts cannot alias one text part.
payload = [p if 'text' in p else {'image_sha256': hashlib.sha256(Path(p['image']).read_bytes()).hexdigest()} for p in parts]
return hashlib.sha256(_json(payload).encode()).hexdigest()
class _CheckedEngine:
def __init__(self, engine):
self.backend = engine
def identity(self):
value = self.backend.identity
return value() if callable(value) else value
def prefill(self, document):
return self.backend.prefill(copy.deepcopy(document))
def ask(self, state, block, n, execution='cached'):
result = self.backend.ask(state, block, n, execution=execution)
logits = np.asarray(result.get('letter_logits'), dtype=float)
tokens = result.get('branch_tokens')
if logits.shape != (n,) or not np.isfinite(logits).all():
raise ValueError('engine returned invalid letter logits')
if type(tokens) is not int or tokens < 0:
raise ValueError('engine returned invalid branch token count')
return result
class Service(ReadoutService):
"""Shared existing readout logic with CUDA-safe persistence and identity checks."""
def __init__(self, store, engine, design=None):
design = copy.deepcopy(design or {'single_choice': 'R', 'ordered': 'R'})
allowed = {'R', 'S', 'P', 'R+avg2', 'R+avg3', 'R+avgall', 'R+debias'}
if any(design.get(k) not in allowed for k in ('single_choice', 'ordered')):
raise ValueError('unsupported readout design')
if design['ordered'] not in {'R', 'S'}:
raise ValueError('ordered design must preserve caller level order (R or S)')
if 'R+debias' in design.values():
prior = design.get('position_prior', {})
for n in range(2, 9):
values = np.asarray(prior.get(str(n+2)), dtype=float)
if values.shape != (n+2,) or not np.isfinite(values).all():
raise ValueError('debias design needs finite priors for 2 to 8 options')
super().__init__(store, _CheckedEngine(engine), design)
self.runtime_identity = _identity(self.engine)
self.design_sha256 = hashlib.sha256(_json(self.design).encode()).hexdigest()
def _check_runtime(self):
if _identity(self.engine) != self.runtime_identity:
self.state = self.state_id = None
raise ValueError('engine identity changed; create a new service instance')
def create(self, document):
parts = _parts(document)
with self.lock:
self._check_runtime()
key = uuid.uuid4().hex
pending = self.store / ('.pending-' + key)
final = self.store / key
pending.mkdir()
try:
saved = []
for i, part in enumerate(parts):
if 'text' in part:
saved.append(dict(part))
else:
source = Path(part['image'])
filename = f'image-{i}{source.suffix}'
(pending / filename).write_bytes(source.read_bytes())
saved.append({'image': str((final / filename).resolve())})
digest_parts = [p if 'text' in p else {'image': str(pending / Path(p['image']).name)} for p in saved]
record = {'schema': 'cuda-service-inputs-v1', 'state_id': key,
'document': saved, 'document_sha256': _digest(digest_parts),
'runtime_identity': self.runtime_identity,
'persistence': 'immutable_inputs_restart_reprefill'}
(pending / 'record.json').write_text(_json(record))
pending.rename(final)
except BaseException:
shutil.rmtree(pending, ignore_errors=True)
raise
return {'state_id': key, 'contract': 'solomon-answer-contract-v3',
'document_sha256': record['document_sha256'],
'persistence': record['persistence']}
def _warm(self, key):
if not isinstance(key, str) or re.fullmatch('[0-9a-f]{32}', key) is None:
raise ValueError('invalid state identifier')
self._check_runtime()
root = self.store / key
record = json.loads((root / 'record.json').read_text())
if record.get('schema') != 'cuda-service-inputs-v1' or record.get('state_id') != key:
raise ValueError('invalid saved state record')
if record.get('runtime_identity') != self.runtime_identity:
raise ValueError('saved state uses a different runtime identity')
parts = _parts(record['document'])
for part in parts:
if 'image' in part and Path(part['image']).resolve().parent != root.resolve():
raise ValueError('saved image must stay inside its state directory')
if _digest(parts) != record.get('document_sha256'):
raise ValueError('saved document content changed')
if self.state_id != key:
# Clear both before prefill: a failed prefill cannot reuse a stale state.
self.state = self.state_id = None
self.state = self.engine.prefill(parts)
self.state_id = key
return self.state
def ask(self, key, task, **kwargs):
if task not in ('boolean', 'single', 'ordered', 'multilabel', 'entity'):
raise ValueError('unknown answer type')
with self.lock:
result = super().ask(key, task, **kwargs)
result.update(runtime_identity=copy.deepcopy(self.runtime_identity),
design_sha256=self.design_sha256,
persistence='immutable_inputs_restart_reprefill')
return result