"""ROS 2 request -> RAG-grounded package, validated against the indexed interface definitions. Two backends generate the files: - "opencode" (default): OpenCode runs headless in a throwaway workspace and writes the package itself; the RAG context is attached to its message. Repairs continue the same OpenCode session. - "direct": one chat call returns the files as fenced blocks. Used when OpenCode is not installed, and in tests. """ from __future__ import annotations import logging import re from dataclasses import dataclass, field from typing import Protocol from rag.chunking import interface_kind from rag.store import Hit, RagStore, format_context from .config import Settings from .llm import LLM, LLMError, Message from .opencode_runner import OpenCodeError, OpenCodeWorkspace from .ros_check import CodeValidation, load_interfaces, validate_package from .routing import is_free log = logging.getLogger("rosdiff.codegen") DISTROS = ("humble", "jazzy") LANGUAGES = ("python", "cpp") TASK_TEMPLATE = """Create a complete, buildable ROS 2 {distro} package in {language_name} for this request: {request} Requirements: - Package name: {package_name}. Put the package in the directory ./{package_name}/. - {layout} - Use only message/service types whose definitions appear in the ROS 2 CONTEXT (they are verified to exist in {distro}). Declare every package you use in package.xml. - Follow the APIs shown in the ROS 2 CONTEXT (rclpy / rclcpp, launch). Include a launch file in ./{package_name}/launch/ that starts the node(s) with sensible parameters. - Include short comments where behaviour is not obvious. No placeholders or TODOs: the code must run as is. """ LAYOUT = { "python": "ament_python layout: package.xml, setup.py, setup.cfg, resource/, /__init__.py, " "/.py with a main() registered under console_scripts, and the launch file installed via data_files.", "cpp": "ament_cmake layout: package.xml, CMakeLists.txt (find_package, add_executable, " "ament_target_dependencies, install targets and launch dir, ament_package()), src/.cpp.", } DIRECT_FORMAT = """Output every file of the package in this exact format and nothing else between files: FILE: ``` ``` """ class CodeBackend(Protocol): name: str def generate(self, task: str, context: str) -> tuple[dict[str, str], str]: ... def repair(self, feedback: str, context: str) -> tuple[dict[str, str], str]: ... @dataclass class CodeAttempt: files: dict[str, str] validation: dict transcript: str model: str = "" usage: dict = field(default_factory=dict) @dataclass class CodeResult: valid: bool files: dict[str, str] validation: CodeValidation attempts: list[CodeAttempt] = field(default_factory=list) sources: list[str] = field(default_factory=list) backend: str = "" model: str = "" # --------------------------------------------------------------------------- retrieval def mentioned_interfaces(request: str, interfaces: dict[str, str]) -> list[str]: """Interfaces named in the request, e.g. 'LaserScan' or 'sensor_msgs/LaserScan'.""" words = set(re.findall(r"[A-Za-z_]+(?:/[A-Za-z_]+)*", request)) found = [] for name in interfaces: pkg, _, base = name.split("/") if base in words or f"{pkg}/{base}" in words or name in words: found.append(name) return sorted(found) def nested_types(definition: str, own_package: str, interfaces: dict[str, str]) -> list[str]: """Field types used inside a message definition, resolved to full interface names.""" out = [] for line in definition.splitlines(): line = line.split("#", 1)[0].strip() if not line or "=" in line.split()[0]: continue type_name = re.sub(r"\[.*\]$", "", line.split()[0]) if "/" in type_name: pkg, base = type_name.split("/")[0], type_name.split("/")[-1] else: pkg, base = own_package, type_name full = f"{pkg}/msg/{base}" if full in interfaces: out.append(full) return out def retrieve(store: RagStore, request: str, distro: str, language: str, interfaces: dict[str, str]) -> list[Hit]: hits: list[Hit] = [] # 1. Exact definitions of interfaces the request names, plus the types they contain (one level). named = mentioned_interfaces(request, interfaces) expanded = list(named) for name in named: for sub in nested_types(interfaces[name], name.split("/")[0], interfaces): if sub not in expanded: expanded.append(sub) for name in expanded[:12]: kind = interface_kind(name) hits.append( Hit( f"[ROS 2 {distro} {kind} {name}]\n{interfaces[name].strip()}", f"interface:{distro}:{name}", 0.0, {"kind": "interface"}, ) ) # 2. Other interfaces that look relevant. hits += store.query("ros2", request, k=4, where={"$and": [{"kind": "interface"}, {"distro": distro}]}) # 3. Tutorials and how-tos for the right distro and language. lang = "Python rclpy" if language == "python" else "C++ rclcpp" hits += store.query("ros2", f"{lang} {request}", k=6, where={"$and": [{"kind": "doc"}, {"distro": distro}]}) seen, unique = set(), [] for h in hits: if h.text not in seen: seen.add(h.text) unique.append(h) return unique # --------------------------------------------------------------------------- backends _FILE_BLOCK = re.compile(r"^FILE:\s*(?P\S+)\s*\n```[\w+-]*\n(?P.*?)^```", re.MULTILINE | re.DOTALL) def parse_file_blocks(text: str) -> dict[str, str]: return {m.group("path").strip("`'\""): m.group("body") for m in _FILE_BLOCK.finditer(text)} class DirectBackend: name = "direct" def __init__(self, llm: LLM): self.llm = llm self.messages: list[Message] = [] self.model: str | None = None # set per attempt by the router self.last_usage: dict = {} def _call(self) -> tuple[dict[str, str], str]: completion = ( self.llm.complete(self.messages, model=self.model) if self.model else self.llm.complete(self.messages) ) self.last_usage = dict(completion.usage or {}) self.messages.append(Message("assistant", completion.text)) return parse_file_blocks(completion.text), completion.text def generate(self, task: str, context: str) -> tuple[dict[str, str], str]: self.messages = [ Message( "system", "You are an expert ROS 2 developer who writes correct, buildable packages.\n" + DIRECT_FORMAT ), Message("user", f"ROS 2 CONTEXT:\n{context}\n\nTASK:\n{task}"), ] return self._call() def repair(self, feedback: str, context: str) -> tuple[dict[str, str], str]: self.messages.append( Message( "user", f"{feedback}\n\nFix these problems and output ALL files of the package again in the same format.", ) ) return self._call() class OpenCodeBackend: name = "opencode" def __init__(self, workspace: OpenCodeWorkspace): self.workspace = workspace self.model: str | None = None # set per attempt by the router self.last_usage: dict = {} def _result(self, run) -> tuple[dict[str, str], str]: self.last_usage = run.usage() if run.error: raise LLMError(f"OpenCode: {run.error}") texts = [e.get("part", {}).get("text", "") for e in run.events if e.get("type") == "text"] return self.workspace.collect(), "\n".join(t for t in texts if t) def generate(self, task: str, context: str) -> tuple[dict[str, str], str]: message = ( task + "\nThe attached file contains the ROS 2 CONTEXT. Write the files directly into the " "current directory using your edit tools; do not run shell commands." ) return self._run(message, f"# ROS 2 CONTEXT\n\n{context}\n") def repair(self, feedback: str, context: str) -> tuple[dict[str, str], str]: return self._run(f"{feedback}\n\nFix these problems by editing the files.", f"# ROS 2 CONTEXT\n\n{context}\n") def _run(self, message: str, context: str) -> tuple[dict[str, str], str]: # free debugging models live under OpenCode's own Zen provider, paid ones under "llm" (OpenRouter) if is_free(self.model): provider, _, model = self.model.partition("/") return self._result(self.workspace.run(message, context, model=model, provider=provider)) return self._result(self.workspace.run(message, context, model=self.model)) # --------------------------------------------------------------------------- pipeline def default_package_name(request: str) -> str: words = [w.lower() for w in re.findall(r"[A-Za-z]+", request) if len(w) > 2][:3] return "_".join(words) or "generated_pkg" def generate_code( request: str, distro: str, language: str, backend: CodeBackend, store: RagStore, settings: Settings, package_name: str | None = None, max_attempts: int | None = None, plan=None, ) -> CodeResult: """`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 distro not in DISTROS: raise ValueError(f"distro must be one of {DISTROS}") if language not in LANGUAGES: raise ValueError(f"language must be one of {LANGUAGES}") catalogue = load_interfaces(str(settings.interfaces_file)) interfaces = catalogue.get(distro, {}) hits = retrieve(store, request, distro, language, interfaces) context = format_context(hits, max_chars=16000) package_name = package_name or default_package_name(request) task = TASK_TEMPLATE.format( distro=distro, language_name="Python (rclpy)" if language == "python" else "C++ (rclcpp)", request=request.strip(), package_name=package_name, layout=LAYOUT[language], ) attempts: list[CodeAttempt] = [] files: dict[str, str] = {} validation = CodeValidation(False, ["no attempt made"]) for i in range(max(1, max_attempts or settings.code_max_attempts)): candidates = [m.id for m in candidates_for_attempt(plan, i)] if plan is not None else [None] for n, model in enumerate(candidates): if model is not None: backend.model = model try: files, transcript = ( backend.generate(task, context) if i == 0 else backend.repair(validation.feedback(), context) ) break except (OpenCodeError, LLMError) as e: if n == len(candidates) - 1: raise LLMError(str(e)) from e log.info("model %s failed (%s); falling back to %s", model, e, candidates[n + 1]) validation = validate_package(files, distro, catalogue) used = getattr(backend, "model", None) or getattr(getattr(backend, "llm", None), "model", "") or "" attempts.append( CodeAttempt(dict(files), validation.to_dict(), transcript, used, dict(getattr(backend, "last_usage", {}))) ) if validation.valid: break model = (attempts[-1].model if attempts else "") or settings.code_llm_model return CodeResult( valid=validation.valid, files=files, validation=validation, attempts=attempts, sources=[h.source for h in hits], backend=backend.name, model=model, )