Spaces:
Paused
Paused
Download api/codegen.py from chandrakiran06/rosdiff: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/spaces/chandrakiran06/rosdiff/resolve/main/api/codegen.py
- Command line
-
hf download hf://spaces/chandrakiran06/rosdiff/api/codegen.py
-
curl -L -o codegen.py https://huggingface.co/spaces/chandrakiran06/rosdiff/resolve/main/api/codegen.py
11.8 kB
| """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/<name>, <name>/__init__.py, " | |
| "<name>/<node>.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/<node>.cpp.", | |
| } | |
| DIRECT_FORMAT = """Output every file of the package in this exact format and nothing else between files: | |
| FILE: <relative/path> | |
| ```<language> | |
| <complete file content> | |
| ``` | |
| """ | |
| 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]: ... | |
| class CodeAttempt: | |
| files: dict[str, str] | |
| validation: dict | |
| transcript: str | |
| model: str = "" | |
| usage: dict = field(default_factory=dict) | |
| 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<path>\S+)\s*\n```[\w+-]*\n(?P<body>.*?)^```", 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, | |
| ) | |