rosdiff / api /codegen.py
Chandra Kiran
Free debugger on OpenCode Zen; separate code and scene model lists
536e671 unverified
Raw History Blame Contribute Delete
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]: ...
@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<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,
)