"""Scenario description -> RAG-grounded MJCF scene, validated by loading it in MuJoCo.""" from __future__ import annotations import json import logging import re from dataclasses import dataclass, field from functools import lru_cache from pathlib import Path from rag.store import Hit, RagStore, format_context from .config import Settings from .llm import LLM, LLMError, Message from .mjcf_check import ValidationResult, extract_mjcf, validate_mjcf from .robots import ROBOTS, robot_dir, robot_facts log = logging.getLogger("rosdiff.simulate") # Elements almost every scene uses; their exact attribute lists are always in the prompt. CORE_ELEMENTS = ( "mujoco", "option", "visual-headlight", "asset-texture", "asset-material", "asset-hfield", "default", "body", "body-geom", "body-joint", "body-freejoint", "body-inertial", "body-site", "body-light", "body-camera", "equality-weld", ) SYSTEM_PROMPT = """You are an expert MuJoCo modeller. You write MJCF scene files for robotics simulation. Rules: - Output exactly one MJCF document in a single ```xml fenced block, then at most three short lines of notes. - The root element is . - {robot_rule} - Use only elements and attributes listed in the MJCF SCHEMA / reference context. Do not invent attributes. - Build the environment from primitive geoms (plane, box, sphere, capsule, cylinder, ellipsoid) or a builtin hfield. Do not reference external mesh/texture/hfield files; builtin textures (builtin="checker"/"gradient"/"flat") are fine. - Give every body and geom you add a descriptive name. Add a floor, a light, and sensible friction/mass values. - Objects meant to be pushed, carried or to fall get a ; fixed terrain and furniture do not. - Angles: the robot file sets ; give euler/axisangle values in radians. """ ROBOT_RULE = ( 'Include the robot with as the first child of . Never copy, redefine, or rename ' "the robot's bodies, joints, actuators, meshes, compiler or keyframes." ) NO_ROBOT_RULE = "There is no robot to include; do not use ." @dataclass class Attempt: raw_output: str mjcf: str | None validation: dict latency_ms: int model: str = "" usage: dict = field(default_factory=dict) @dataclass class SimulateResult: valid: bool mjcf: str | None validation: ValidationResult attempts: list[Attempt] = field(default_factory=list) sources: list[str] = field(default_factory=list) model: str = "" @lru_cache def load_schema(path: str) -> dict: p = Path(path) return json.loads(p.read_text()) if p.is_file() else {} def schema_lines(schema: dict, keys) -> str: rows = [] for key in keys: entry = schema.get(key) if entry and entry["attributes"]: rows.append(f"<{entry['element']}> ({entry['path']}): {', '.join(entry['attributes'])}") return "\n".join(rows) def schema_hint(schema: dict, error: str) -> str: """For 'unrecognized attribute ... Element 'geom'' style errors: the element's valid attributes.""" m = re.search(r"Element '(\w+)'", error) if not m: return "" element = m.group(1) keys = [k for k, e in schema.items() if e["element"] == element] return schema_lines(schema, keys) def retrieve(store: RagStore, description: str) -> list[Hit]: """Scenario-specific sections of the MuJoCo docs (the robot and the schema are pinned separately).""" return store.query("mujoco", description, k=6, where={"kind": "doc"}) def robot_example(settings: Settings, robot: str) -> str: """Menagerie's own scene.xml for the robot: the canonical way to include it.""" path = robot_dir(settings.menagerie_dir, ROBOTS[robot]) / "scene.xml" return path.read_text() if path.is_file() else "" def build_messages( description: str, robot: str | None, context: str, facts: str, example: str = "", schema: str = "" ) -> list[Message]: robot_rule = ROBOT_RULE.format(file=ROBOTS[robot].model_file) if robot else NO_ROBOT_RULE user = [] if facts: user.append(f"ROBOT FACTS (read from the model file):\n{facts}") if example: user.append(f"MENAGERIE EXAMPLE (scene.xml shipped with the robot):\n```xml\n{example.strip()}\n```") if schema: user.append(f"MJCF SCHEMA (the only valid attributes of common elements):\n{schema}") user.append(f"REFERENCE CONTEXT (MuJoCo documentation):\n{context}") user.append(f"SCENARIO:\n{description.strip()}") user.append("Write the MJCF scene now.") return [ Message("system", SYSTEM_PROMPT.format(robot_rule=robot_rule)), Message("user", "\n\n".join(user)), ] def generate_scene( description: str, robot: str | None, llm: LLM, store: RagStore, settings: Settings, max_attempts: int | None = None, plan=None, ) -> SimulateResult: """`plan` (api.routing.Plan) picks the model per attempt: the planned tier first, one tier up per repair.""" from .routing import candidates_for_attempt if robot is not None and robot not in ROBOTS: raise ValueError(f"unknown robot {robot!r}; available: {sorted(ROBOTS)}") hits = retrieve(store, description) facts = robot_facts(str(settings.menagerie_dir), robot) if robot else "" example = robot_example(settings, robot) if robot else "" schema = load_schema(str(settings.mjcf_schema_file)) messages = build_messages( description, robot, format_context(hits, max_chars=10000), facts, example, schema_lines(schema, CORE_ELEMENTS) ) sources = ([f"mujoco_menagerie:{ROBOTS[robot].directory}/scene.xml"] if example else []) + [h.source for h in hits] attempts: list[Attempt] = [] validation = ValidationResult(False, "extract", "no attempt made") mjcf = None for i in range(max(1, max_attempts or settings.mjcf_max_attempts)): completion, model = _complete(llm, messages, [m.id for m in candidates_for_attempt(plan, i)] if plan else []) mjcf = extract_mjcf(completion.text) if mjcf is None: validation = ValidationResult(False, "extract", "the response did not contain an MJCF document") else: validation = validate_mjcf( mjcf, robot, settings.menagerie_dir, settings.mjcf_sim_seconds, settings.validate_timeout ) attempts.append( Attempt( completion.text, mjcf, validation.to_dict(), completion.latency_ms, model or completion.model, dict(completion.usage or {}), ) ) if validation.valid: break # Repair: show the model its own output and exactly what MuJoCo said. messages = messages + [ Message("assistant", completion.text), Message( "user", f"{validation.feedback()}{_hint_block(schema, validation.error)}\n\n" "Fix the problem and output the complete corrected MJCF document in one ```xml block. " "Keep everything else the same.", ), ] return SimulateResult( valid=validation.valid, mjcf=mjcf, validation=validation, attempts=attempts, sources=sources, model=attempts[-1].model if attempts else llm.model, ) def _complete(llm: LLM, messages, candidates: list[str]): """Try each candidate model in turn (a free debugger first, then its paid fallback); raise the last error.""" if not candidates: return llm.complete(messages), None for n, model in enumerate(candidates): try: return llm.complete(messages, model=model), model except LLMError as e: if n == len(candidates) - 1: raise log.info("model %s failed (%s); falling back to %s", model, e, candidates[n + 1]) raise AssertionError("unreachable") def _hint_block(schema: dict, error: str) -> str: hint = schema_hint(schema, error) return f"\nValid attributes for that element:\n{hint}" if hint else ""