diff --git a/src/frontend/cua_s1.py b/src/frontend/cua_s1.py new file mode 100644 index 0000000..bde882c --- /dev/null +++ b/src/frontend/cua_s1.py @@ -0,0 +1,136 @@ +"""Small loopback HTTP worker; the Rust frontend remains the public serving layer.""" + +from __future__ import annotations + +import argparse +import json +import logging +import socket +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +from models.cua_s1.multimodal.protocol import ( + MAX_BODY, + InvalidRequest, + MalformedJSON, + decode_request, + parse_request, +) + +LOG = logging.getLogger(__name__) + + +class WorkerServer(ThreadingHTTPServer): + daemon_threads = True + + def __init__(self, address, engine): + self.engine = engine + self.inference_lock = threading.Lock() + super().__init__(address, Handler) + + +class Handler(BaseHTTPRequestHandler): + def setup(self): + super().setup() + self.connection.settimeout(15) + + def log_message(self, format, *args): + # Do not log paths, input images, instructions or arbitrary request headers. + pass + + def send_json(self, status, value): + raw = json.dumps(value, ensure_ascii=False, allow_nan=False).encode("utf-8") + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(raw))) + self.end_headers() + try: + self.wfile.write(raw) + except (BrokenPipeError, ConnectionResetError): + pass + + def do_GET(self): + if self.path == "/health": + self.send_json(200, {"status": "ready", "modality": "multimodal"}) + else: + self.send_json(404, {"detail": "unknown route"}) + + def do_POST(self): + if self.path != "/v1/systemone": + self.send_json(404, {"detail": "unknown route"}) + return + if self.headers.get("Transfer-Encoding"): + self.send_json( + 411, + { + "detail": "Content-Length is required; chunked requests are unsupported" + }, + ) + return + lengths = self.headers.get_all("Content-Length", []) + if len(lengths) != 1: + self.send_json(411, {"detail": "one Content-Length is required"}) + return + try: + length = int(lengths[0]) + except ValueError: + self.send_json(400, {"detail": "invalid Content-Length"}) + return + if length < 0 or length > MAX_BODY: + self.send_json(413, {"detail": "request exceeds body limit"}) + return + if self.headers.get_content_type() != "application/json": + self.send_json(415, {"detail": "Content-Type must be application/json"}) + return + if not self.server.inference_lock.acquire(blocking=False): + self.send_json(503, {"detail": "worker busy"}) + return + try: + raw = self.rfile.read(length) + if len(raw) != length: + self.send_json(400, {"detail": "incomplete body"}) + return + parsed = parse_request(decode_request(raw)) + result = self.server.engine.predict(parsed) + self.send_json(200, result) + except MalformedJSON as exc: + self.send_json(400, {"detail": str(exc)}) + except InvalidRequest as exc: + self.send_json(422, {"detail": str(exc)}) + except (TimeoutError, socket.timeout): + self.send_json(408, {"detail": "request body timed out"}) + except Exception as exc: + LOG.error("inference failed: %s", type(exc).__name__) + self.send_json(500, {"detail": "inference failed"}) + finally: + self.server.inference_lock.release() + + +def main(): + from models.cua_s1.multimodal.model import MultimodalEngine + + p = argparse.ArgumentParser(description=__doc__) + p.add_argument( + "--base", required=True, help="verified local base checkpoint directory" + ) + p.add_argument( + "--adapter", required=True, help="verified local multimodal adapter directory" + ) + p.add_argument("--port", type=int, default=8000) + args = p.parse_args() + logging.basicConfig(level=logging.INFO) + engine = MultimodalEngine(args.base, args.adapter) + engine.warmup() + # Bind only after model loading and a representative inference succeed. + server = WorkerServer(("127.0.0.1", args.port), engine) + LOG.info("multimodal worker ready on 127.0.0.1:%s", args.port) + try: + server.serve_forever() + except KeyboardInterrupt: + pass + finally: + server.server_close() + + +if __name__ == "__main__": + main() diff --git a/src/models/cua_s1/multimodal/model.py b/src/models/cua_s1/multimodal/model.py new file mode 100644 index 0000000..e3655de --- /dev/null +++ b/src/models/cua_s1/multimodal/model.py @@ -0,0 +1,174 @@ +"""Direct Transformers/PEFT execution. No production dependency on cua_s1.""" + +from __future__ import annotations + +import hashlib +import json +from pathlib import Path + +from .protocol import InvalidRequest, Question, Request, answer, build_messages + +REFERENCE_REVISION = "0e75660ce4c2edda519e0c795fa3ad98abf4e76f" +BASE_REVISION = "851bf6e806efd8d0a36b00ddf55e13ccb7b8cd0a" +ADAPTER_REVISION = "16818868b0cc7813808aae4e87b417657046ab79" +IDENTITY = f"cua-ai/cua-s1-4b-0.2@{ADAPTER_REVISION}:multimodal" +WEIGHTS_MANIFEST_SHA256 = ( + "9820bd232c5762f114e19680c0f8203d7e1faaf8a60c196cfe01964d6d8a6c09" +) +MAX_TOKENS = 4096 + + +def letter_ids(tokenizer, count: int) -> list[int]: + ids = [] + for index in range(count): + encoded = tokenizer.encode(chr(65 + index), add_special_tokens=False) + if len(encoded) != 1: + raise ValueError("each candidate letter must be a single token") + ids.append(encoded[0]) + return ids + + +def validate_adapter_config(config: dict): + targets = { + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", + "linear_fc1", + "linear_fc2", + } + if ( + config.get("peft_type") != "LORA" + or config.get("r") != 16 + or config.get("lora_alpha") != 32 + or set(config.get("target_modules", [])) != targets + or config.get("base_model_name_or_path") != "Qwen/Qwen3.5-4B" + ): + raise ValueError("expected the pinned 0.2 multimodal LoRA adapter") + + +def parse_weights_manifest(raw: bytes) -> dict: + """Accept only the manifest from the pinned upstream reference commit.""" + if hashlib.sha256(raw).hexdigest() != WEIGHTS_MANIFEST_SHA256: + raise ValueError("upstream weights manifest checksum mismatch") + return json.loads(raw) + + +def verify_weights(base: Path, adapter: Path): + """Check local artifacts before assigning the pinned identity to responses.""" + lock = parse_weights_manifest((base.parent / "weights.lock.json").read_bytes()) + allowed = {base: set(), adapter: set()} + for artifact in lock["artifacts"]: + for name, expected in artifact["files"].items(): + if artifact["role"] == "adapter": + if not name.startswith("multimodal/"): + continue + path = adapter / name.removeprefix("multimodal/") + else: + path = base / name + root = adapter if artifact["role"] == "adapter" else base + allowed[root].add(path.relative_to(root).as_posix()) + if not path.is_file() or path.stat().st_size != expected["size"]: + raise ValueError(f"missing or wrong-size pinned artifact: {path.name}") + with path.open("rb") as handle: + digest = hashlib.file_digest(handle, "sha256").hexdigest() + if digest != expected["sha256"]: + raise ValueError(f"checksum mismatch: {path.name}") + for root, names in allowed.items(): + for path in root.rglob("*"): + relative = path.relative_to(root) + if path.is_file() and relative.parts[0] != ".cache": + if relative.as_posix() not in names: + raise ValueError( + f"unlisted artifact may override pinned files: {relative}" + ) + + +class MultimodalEngine: + def __init__( + self, base: str, adapter: str, device: str = "cuda", dtype: str = "bfloat16" + ): + import torch + from peft import PeftModel + from peft.tuners.lora import LoraLayer + from transformers import ( + AutoModelForImageTextToText, + AutoProcessor, + AutoTokenizer, + ) + + base_path, adapter_path = Path(base), Path(adapter) + verify_weights(base_path, adapter_path) + validate_adapter_config( + json.loads((adapter_path / "adapter_config.json").read_text()) + ) + self.tokenizer = AutoTokenizer.from_pretrained(base, local_files_only=True) + self.processor = AutoProcessor.from_pretrained(base, local_files_only=True) + model = AutoModelForImageTextToText.from_pretrained( + base, + torch_dtype=getattr(torch, dtype), + device_map=device, + local_files_only=True, + ) + self.model = PeftModel.from_pretrained(model, adapter, local_files_only=True) + modules = [ + name + for name, module in self.model.named_modules() + if isinstance(module, LoraLayer) + ] + if len(modules) != 178 or not any(".visual." in name for name in modules): + raise RuntimeError( + "multimodal adapter did not attach to all 178 expected modules" + ) + self.adapter_modules = len(modules) + self.model.eval() + self.dtype = dtype + + def prepare(self, image, question: Question): + messages = build_messages(question) + text = self.processor.apply_chat_template( + messages, tokenize=False, add_generation_prompt=True + ) + inputs = self.processor(text=[text], images=[image], return_tensors="pt") + if inputs["input_ids"].shape[-1] > MAX_TOKENS: + raise InvalidRequest(f"processed prompt exceeds {MAX_TOKENS} tokens") + return inputs + + def score(self, inputs, question: Question) -> list[float]: + import torch + + ids = letter_ids(self.tokenizer, len(question.keys)) + inputs = inputs.to(self.model.device) + with torch.no_grad(): + output = self.model(**inputs) + logits = output.logits[0, -1, :] + return torch.softmax( + logits[torch.tensor(ids, device=logits.device)].float(), dim=-1 + ).tolist() + + def predict(self, request: Request) -> dict: + # Validate all processed lengths before executing any question. + prepared = [self.prepare(request.image, q) for q in request.questions] + answers = { + q.name: answer(q, self.score(inputs, q)) + for q, inputs in zip(request.questions, prepared) + } + return { + "model": IDENTITY, + "answers": answers, + "usage": { + "input_tokens": sum(x["input_ids"].shape[-1] for x in prepared), + "output_tokens": 0, + }, + } + + def warmup(self): + from PIL import Image + + q = Question( + "warmup", ("continue", "cancel"), ("Continue", "Cancel"), "Continue" + ) + self.predict(Request(Image.new("RGB", (224, 224), "white"), (q,))) diff --git a/src/models/cua_s1/multimodal/protocol.py b/src/models/cua_s1/multimodal/protocol.py new file mode 100644 index 0000000..4346257 --- /dev/null +++ b/src/models/cua_s1/multimodal/protocol.py @@ -0,0 +1,261 @@ +"""Bounded screenshot-only extension of the Cua-S1 choice contract in PR #11.""" + +from __future__ import annotations + +import base64 +import binascii +import io +import json +import math +from dataclasses import dataclass + +from PIL import Image, UnidentifiedImageError + +MODEL = "cua-s1-4b-0.2" +MAX_BODY = 8 * 1024 * 1024 +MAX_IMAGE_BYTES = 4 * 1024 * 1024 +MAX_PIXELS = 1024 * 1024 +MAX_SIDE = 2048 +MAX_ASPECT_RATIO = 200 # Pinned Transformers Qwen image processor's smart_resize. +MAX_QUESTIONS = 8 +MAX_TEXT = 16384 + +# Prompt and letter layout follow trycua/cua at 0e75660ce4c2edda519e0c795fa3ad98abf4e76f. +# MIT License +# +# Copyright (c) 2025 Cua AI, Inc. +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +SYSTEM_PROMPT = ( + "You are a one-pass computer-use decision model. You are shown the " + "current state of a screen and a fixed, closed list of candidate " + "(element, action) options, each given a single letter. Choose exactly " + "one option: the single best next action to take. Answer with ONLY that " + "option's letter -- no words, no punctuation, no explanation." +) + + +class InvalidRequest(ValueError): + """Input cannot be evaluated under the supported contract.""" + + +class MalformedJSON(InvalidRequest): + """The body is not a usable UTF-8 JSON object (HTTP 400).""" + + +@dataclass(frozen=True) +class Question: + name: str + keys: tuple[str, ...] + labels: tuple[str, ...] + goal: str + + +@dataclass(frozen=True) +class Request: + image: Image.Image + questions: tuple[Question, ...] + + +def _object(pairs): + result = {} + for key, value in pairs: + if key in result: + raise MalformedJSON("duplicate JSON keys are not supported") + result[key] = value + return result + + +def _nonfinite(value): + raise MalformedJSON("non-finite JSON numbers are not supported") + + +def decode_request(raw: bytes) -> dict: + if len(raw) > MAX_BODY: + raise InvalidRequest("request body exceeds 8 MiB") + try: + value = json.loads( + raw.decode("utf-8"), object_pairs_hook=_object, parse_constant=_nonfinite + ) + # The decoder accepts lone surrogate escapes and overflowing floats. + # Reject both before any text can reach the tokenizer or a response. + json.dumps(value, ensure_ascii=False, allow_nan=False).encode("utf-8") + except MalformedJSON: + raise + except (ValueError, UnicodeError, RecursionError) as exc: + raise MalformedJSON( + "request body must contain valid JSON and UTF-8 text" + ) from exc + if not isinstance(value, dict): + raise MalformedJSON("request must be a JSON object") + return value + + +def _text(value, field): + if not isinstance(value, (str, dict, list)): + raise InvalidRequest(f"{field} must be a string, object or array") + try: + result = ( + value + if isinstance(value, str) + else json.dumps(value, ensure_ascii=False, allow_nan=False) + ) + except (ValueError, TypeError, RecursionError) as exc: + raise InvalidRequest(f"invalid {field}") from exc + if len(result) > MAX_TEXT: + raise InvalidRequest(f"{field} exceeds {MAX_TEXT} characters") + if any( + token in result + for token in ( + "<|image_pad|>", + "<|video_pad|>", + "<|vision_start|>", + "<|vision_end|>", + ) + ): + raise InvalidRequest(f"{field} contains an unsupported media control token") + return result + + +def _image(state): + if not isinstance(state, dict) or set(state) != {"image"}: + raise InvalidRequest("state must contain exactly one image data URL") + url = state["image"] + if not isinstance(url, str): + raise InvalidRequest("state.image must be a PNG/JPEG base64 data URL") + prefix, separator, encoded = url.partition(",") + expected = {"data:image/png;base64": "PNG", "data:image/jpeg;base64": "JPEG"} + if not separator or prefix not in expected: + raise InvalidRequest("only inline PNG/JPEG images are supported") + if len(encoded) > 4 * ((MAX_IMAGE_BYTES + 2) // 3): + raise InvalidRequest("encoded image exceeds 4 MiB") + try: + raw = base64.b64decode(encoded, validate=True) + if len(raw) > MAX_IMAGE_BYTES: + raise InvalidRequest("image exceeds 4 MiB") + with Image.open(io.BytesIO(raw)) as source: + if source.format != expected[prefix]: + raise InvalidRequest("image format does not match its MIME type") + w, h = source.size + if max(w, h) > MAX_ASPECT_RATIO * min(w, h): + raise InvalidRequest( + f"image aspect ratio must not exceed {MAX_ASPECT_RATIO}:1" + ) + if ( + max(w, h) > MAX_SIDE + or w * h > MAX_PIXELS + or getattr(source, "n_frames", 1) != 1 + ): + raise InvalidRequest( + "image must be single-frame, at most 2048 per side and 1048576 pixels" + ) + source.load() + return source.convert("RGB") + except ( + binascii.Error, + UnidentifiedImageError, + OSError, + Image.DecompressionBombError, + ValueError, + ) as exc: + if isinstance(exc, InvalidRequest): + raise + raise InvalidRequest("invalid image data") from exc + + +def parse_request(value: dict) -> Request: + if not isinstance(value, dict) or set(value) != {"model", "state", "questions"}: + raise InvalidRequest("request must contain model, state and questions only") + if value["model"] != MODEL: + raise InvalidRequest(f"model must be {MODEL}") + questions = value["questions"] + if not isinstance(questions, dict) or not 1 <= len(questions) <= MAX_QUESTIONS: + raise InvalidRequest("questions must contain 1 to 8 questions") + parsed = [] + for name, q in questions.items(): + if not isinstance(name, str) or not name or len(name) > 256: + raise InvalidRequest("question names must contain 1 to 256 characters") + if not isinstance(q, dict) or q.get("type") != "choice": + raise InvalidRequest("only choice questions are supported") + if set(q) - {"type", "instructions", "criteria"}: + raise InvalidRequest("unsupported question fields") + if "instructions" not in q: + raise InvalidRequest("instructions is required for every question") + criteria = q.get("criteria") + if not isinstance(criteria, dict) or not 1 <= len(criteria) <= 26: + raise InvalidRequest("choice requires 1 to 26 options") + labels = [] + for key, label in criteria.items(): + if not isinstance(key, str) or not key or len(key) > 256: + raise InvalidRequest("option keys must contain 1 to 256 characters") + text = _text(key if label is None else label, "criteria") + labels.append(json.dumps(text, ensure_ascii=False)[1:-1]) + goal = ( + "" + if q["instructions"] is None + else _text(q["instructions"], "instructions") + ) + if len(goal) + sum(map(len, labels)) > MAX_TEXT: + raise InvalidRequest("combined question text exceeds 16384 characters") + parsed.append(Question(name, tuple(criteria), tuple(labels), goal)) + return Request(_image(value["state"]), tuple(parsed)) + + +def build_messages(q: Question, image_marker: str = "inline.png") -> list[dict]: + lines = "\n".join( + f'{chr(65 + i)}. Decision "{label}" -> select' + for i, label in enumerate(q.labels) + ) + text = ( + (f"Goal: {q.goal}\n\n" if q.goal else "") + + "App: Cua Driver\nTask family: closed-candidate decision\n\n" + + "The current screenshot is attached.\n\n" + + f"Options:\n{lines}\n\nAnswer with a single letter." + ) + return [ + {"role": "system", "content": SYSTEM_PROMPT}, + { + "role": "user", + "content": [ + {"type": "image", "image": image_marker}, + {"type": "text", "text": text}, + ], + }, + ] + + +def answer(q: Question, probabilities: list[float]) -> dict: + if len(probabilities) != len(q.keys) or any( + not math.isfinite(p) or not 0 <= p <= 1 for p in probabilities + ): + raise ValueError("model returned invalid probabilities") + if not math.isclose(sum(probabilities), 1.0, abs_tol=1e-5): + raise ValueError("model probabilities do not sum to one") + n = len(probabilities) + confidence = ( + 1.0 + if n == 1 + else 1 + sum(p * math.log(p) for p in probabilities if p) / math.log(n) + ) + return { + "type": "choice", + "choice": q.keys[max(range(n), key=probabilities.__getitem__)], + "probabilities": dict(zip(q.keys, probabilities)), + "confidence": max(0.0, min(1.0, confidence)), + }