mac-compute-space / app /models.py
josephrw's picture
Upload folder using huggingface_hub
c4916b2 verified
Raw
History Blame Contribute Delete
4.6 kB
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime
from enum import Enum
from typing import Any, Optional, List, Dict
class WorkerRuntimeType(str, Enum):
SAFARI_WASM = "safari_wasm"
SAFARI_WEBGPU = "safari_webgpu"
NATIVE_COREML = "native_coreml"
NATIVE_MLX = "native_mlx"
NATIVE_LLAMACPP = "native_llamacpp"
class JobType(str, Enum):
TEXT_EMBEDDING = "text_embedding"
IMAGE_CLASSIFICATION = "image_classification"
IMAGE_EMBEDDING = "image_embedding"
LOCAL_OCR = "local_ocr"
AUDIO_TRANSCRIPTION = "audio_transcription"
SMALL_LLM_GENERATE = "small_llm_generate"
PRIVACY_REDACTION = "privacy_redaction"
SENSOR_CLASSIFICATION = "sensor_classification"
class JobStatus(str, Enum):
QUEUED = "queued"
ASSIGNED = "assigned"
RUNNING = "running"
COMPLETED = "completed"
FAILED = "failed"
EXPIRED = "expired"
REJECTED = "rejected"
class PrivacyMode(str, Enum):
RAW_INPUT_REMOTE = "raw_input_remote"
HASH_ONLY_REMOTE = "hash_only_remote"
LOCAL_ONLY_RESULT_ONLY = "local_only_result_only"
LOCAL_ONLY_REDACTED_RESULT = "local_only_redacted_result"
@dataclass
class DeviceCapability:
capability_name: str
runtime_type: WorkerRuntimeType
model_id: Optional[str] = None
model_hash: Optional[str] = None
quantization: Optional[str] = None
max_input_bytes: Optional[int] = None
estimated_latency_ms: Optional[int] = None
@dataclass
class WorkerState:
worker_id: str
session_id: str
runtime_type: WorkerRuntimeType
capabilities: List[DeviceCapability] = field(default_factory=list)
battery_level: Optional[float] = None
thermal_state: Optional[str] = None
network_type: Optional[str] = None
device_label_hash: Optional[str] = None
device_public_key: Optional[str] = None
joined_at: datetime = field(default_factory=datetime.utcnow)
last_heartbeat: datetime = field(default_factory=datetime.utcnow)
jobs_completed: int = 0
status: str = "active"
@dataclass
class InferenceJob:
job_id: str
session_id: str
job_type: JobType
privacy_mode: PrivacyMode
payload: Dict[str, Any]
status: JobStatus = JobStatus.QUEUED
constraints: Dict[str, Any] = field(default_factory=dict)
created_at: datetime = field(default_factory=datetime.utcnow)
assigned_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
worker_id: Optional[str] = None
result: Optional[Dict[str, Any]] = None
error_reason: Optional[str] = None
@dataclass
class JobResult:
job_id: str
worker_id: str
output: Dict[str, Any]
latency_ms: int
input_hash: Optional[str] = None
output_hash: Optional[str] = None
device_signature: Optional[str] = None
privacy_mode: PrivacyMode = PrivacyMode.RAW_INPUT_REMOTE
@dataclass
class ComputeReceipt:
receipt_id: str
session_id: str
job_id: str
worker_id: str
capability: str
runtime_type: WorkerRuntimeType
job_type: JobType
privacy_mode: PrivacyMode
latency_ms: int
started_at: datetime
finished_at: datetime
model_id: Optional[str] = None
model_hash: Optional[str] = None
input_hash: Optional[str] = None
output_hash: Optional[str] = None
device_public_key: Optional[str] = None
device_signature: Optional[str] = None
server_signature: Optional[str] = None
receipt_hash: Optional[str] = None
previous_receipt_hash: Optional[str] = None
@dataclass
class ReceiptVerificationResult:
valid: bool
receipt_id: str
hash_match: bool
device_signature_valid: Optional[bool] = None
server_signature_valid: Optional[bool] = None
chain_valid: Optional[bool] = None
error: Optional[str] = None
@dataclass
class DevicePolicy:
session_id: str
allowed_job_types: List[JobType] = field(default_factory=list)
denied_job_types: List[JobType] = field(default_factory=list)
require_battery_above: float = 0.2
require_thermal_nominal: bool = True
max_concurrent_jobs: int = 1
privacy_default: PrivacyMode = PrivacyMode.RAW_INPUT_REMOTE
@dataclass
class SessionState:
session_id: str
owner_label: Optional[str] = None
secret: str = ""
created_at: datetime = field(default_factory=datetime.utcnow)
expires_at: datetime = field(default_factory=datetime.utcnow)
workers: Dict[str, WorkerState] = field(default_factory=dict)
policy: DevicePolicy = field(default_factory=lambda: DevicePolicy(session_id=""))
status: str = "active"