Spaces:
Paused
Paused
File size: 11,812 Bytes
bd0a5c7 da839c2 bd0a5c7 536e671 bd0a5c7 225dc91 bd0a5c7 536e671 bd0a5c7 5b063c2 bd0a5c7 225dc91 bd0a5c7 5b063c2 bd0a5c7 5b063c2 bd0a5c7 5b063c2 bd0a5c7 5b063c2 bd0a5c7 536e671 bd0a5c7 536e671 bd0a5c7 5b063c2 bd0a5c7 5b063c2 536e671 5b063c2 bd0a5c7 536e671 bd0a5c7 5b063c2 bd0a5c7 5b063c2 bd0a5c7 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 | """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,
)
|