ferrum-cli 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ferrum/__init__.py +6 -0
- ferrum/agent.py +387 -0
- ferrum/cli.py +479 -0
- ferrum/config.py +178 -0
- ferrum/context.py +379 -0
- ferrum/model.py +302 -0
- ferrum/patch.py +65 -0
- ferrum/prompts/system.md +89 -0
- ferrum/safety.py +124 -0
- ferrum/tools.py +296 -0
- ferrum/verifier.py +227 -0
- ferrum_cli-0.1.0.dist-info/METADATA +74 -0
- ferrum_cli-0.1.0.dist-info/RECORD +17 -0
- ferrum_cli-0.1.0.dist-info/WHEEL +5 -0
- ferrum_cli-0.1.0.dist-info/entry_points.txt +2 -0
- ferrum_cli-0.1.0.dist-info/licenses/LICENSE +211 -0
- ferrum_cli-0.1.0.dist-info/top_level.txt +1 -0
ferrum/context.py
ADDED
|
@@ -0,0 +1,379 @@
|
|
|
1
|
+
"""Project discovery: find the root, walk it, summarise what's there."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
import os
|
|
7
|
+
import re
|
|
8
|
+
from dataclasses import dataclass, field
|
|
9
|
+
from enum import Enum
|
|
10
|
+
from pathlib import Path, PurePosixPath
|
|
11
|
+
|
|
12
|
+
from ferrum.config import Config
|
|
13
|
+
from ferrum.safety import DENIED_DIRS, is_denied, is_within
|
|
14
|
+
|
|
15
|
+
log = logging.getLogger(__name__)
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class Language(str, Enum):
|
|
19
|
+
C = "c"
|
|
20
|
+
CPP = "cpp"
|
|
21
|
+
RUST = "rust"
|
|
22
|
+
ASM = "asm"
|
|
23
|
+
PYTHON = "python"
|
|
24
|
+
|
|
25
|
+
@property
|
|
26
|
+
def label(self) -> str:
|
|
27
|
+
return {
|
|
28
|
+
Language.C: "C",
|
|
29
|
+
Language.CPP: "C++",
|
|
30
|
+
Language.RUST: "Rust",
|
|
31
|
+
Language.ASM: "Assembly",
|
|
32
|
+
Language.PYTHON: "Python",
|
|
33
|
+
}[self]
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class BuildSystem(str, Enum):
|
|
37
|
+
MAKE = "make"
|
|
38
|
+
CMAKE = "cmake"
|
|
39
|
+
CARGO = "cargo"
|
|
40
|
+
PYTEST = "pytest"
|
|
41
|
+
|
|
42
|
+
@property
|
|
43
|
+
def label(self) -> str:
|
|
44
|
+
return {
|
|
45
|
+
BuildSystem.MAKE: "Make",
|
|
46
|
+
BuildSystem.CMAKE: "CMake",
|
|
47
|
+
BuildSystem.CARGO: "Cargo",
|
|
48
|
+
BuildSystem.PYTEST: "pytest",
|
|
49
|
+
}[self]
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
SOURCE_EXTENSIONS: dict[str, Language] = {
|
|
53
|
+
".c": Language.C,
|
|
54
|
+
".h": Language.C,
|
|
55
|
+
".cc": Language.CPP,
|
|
56
|
+
".cpp": Language.CPP,
|
|
57
|
+
".cxx": Language.CPP,
|
|
58
|
+
".c++": Language.CPP,
|
|
59
|
+
".hh": Language.CPP,
|
|
60
|
+
".hpp": Language.CPP,
|
|
61
|
+
".hxx": Language.CPP,
|
|
62
|
+
".rs": Language.RUST,
|
|
63
|
+
".s": Language.ASM,
|
|
64
|
+
".asm": Language.ASM,
|
|
65
|
+
".py": Language.PYTHON,
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
MAKE_NAMES = frozenset({"makefile", "gnumakefile"})
|
|
69
|
+
CMAKE_NAMES = frozenset({"cmakelists.txt"})
|
|
70
|
+
CARGO_NAMES = frozenset({"cargo.toml"})
|
|
71
|
+
PYTEST_NAMES = frozenset({"pyproject.toml", "setup.py", "pytest.ini", "tox.ini"})
|
|
72
|
+
BUILD_FILE_NAMES = MAKE_NAMES | CMAKE_NAMES | CARGO_NAMES | PYTEST_NAMES
|
|
73
|
+
|
|
74
|
+
PROJECT_MARKERS = (
|
|
75
|
+
".git",
|
|
76
|
+
"cargo.toml",
|
|
77
|
+
"cmakelists.txt",
|
|
78
|
+
"makefile",
|
|
79
|
+
"gnumakefile",
|
|
80
|
+
"pyproject.toml",
|
|
81
|
+
"setup.py",
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
@dataclass(frozen=True)
|
|
86
|
+
class FileEntry:
|
|
87
|
+
"""A file we might show the model, path relative to the project root."""
|
|
88
|
+
|
|
89
|
+
path: str
|
|
90
|
+
size: int = 0
|
|
91
|
+
language: Language | None = None
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
@dataclass
|
|
95
|
+
class ProjectContext:
|
|
96
|
+
"""What we know about the project after a quick look."""
|
|
97
|
+
|
|
98
|
+
root: Path
|
|
99
|
+
languages: list[Language] = field(default_factory=list)
|
|
100
|
+
build_systems: list[BuildSystem] = field(default_factory=list)
|
|
101
|
+
files: list[FileEntry] = field(default_factory=list)
|
|
102
|
+
name: str = ""
|
|
103
|
+
|
|
104
|
+
def __post_init__(self) -> None:
|
|
105
|
+
if not self.name:
|
|
106
|
+
self.name = self.root.name or str(self.root)
|
|
107
|
+
|
|
108
|
+
def display_path(self) -> str:
|
|
109
|
+
"""Project path relative to where the user is standing, usually."""
|
|
110
|
+
try:
|
|
111
|
+
rel = os.path.relpath(self.root, Path.cwd())
|
|
112
|
+
except ValueError:
|
|
113
|
+
return str(self.root)
|
|
114
|
+
rel = rel.replace("\\", "/")
|
|
115
|
+
if rel == ".":
|
|
116
|
+
return "."
|
|
117
|
+
if not rel.startswith("."):
|
|
118
|
+
rel = "./" + rel
|
|
119
|
+
return rel
|
|
120
|
+
|
|
121
|
+
def summary(self) -> str:
|
|
122
|
+
"""The block the CLI prints after `ferrum .`."""
|
|
123
|
+
languages = ", ".join(lang.label for lang in self.languages) or "unknown"
|
|
124
|
+
builds = ", ".join(b.label for b in self.build_systems) or "unknown"
|
|
125
|
+
|
|
126
|
+
lines = [
|
|
127
|
+
"Ferrum",
|
|
128
|
+
"─" * 6,
|
|
129
|
+
"",
|
|
130
|
+
f"Project: {self.display_path()}",
|
|
131
|
+
f"Languages: {languages}",
|
|
132
|
+
f"Build system: {builds}",
|
|
133
|
+
"",
|
|
134
|
+
"Files:",
|
|
135
|
+
]
|
|
136
|
+
if self.files:
|
|
137
|
+
lines.extend(f" {entry.path}" for entry in self.files)
|
|
138
|
+
else:
|
|
139
|
+
lines.append(" (none)")
|
|
140
|
+
return "\n".join(lines)
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
def find_project_root(start: Path) -> Path:
|
|
144
|
+
"""Nearest directory at or above start holding a project marker."""
|
|
145
|
+
start = Path(start).resolve()
|
|
146
|
+
for candidate in (start, *start.parents):
|
|
147
|
+
if any((candidate / marker).exists() for marker in PROJECT_MARKERS):
|
|
148
|
+
return candidate
|
|
149
|
+
return start
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
@dataclass(frozen=True)
|
|
153
|
+
class _Rule:
|
|
154
|
+
prefix: str # directory the .gitignore lived in, "" for the root one
|
|
155
|
+
regex: re.Pattern[str]
|
|
156
|
+
negated: bool
|
|
157
|
+
dir_only: bool
|
|
158
|
+
anchored: bool
|
|
159
|
+
|
|
160
|
+
def matches(self, rel: str, is_dir: bool) -> bool:
|
|
161
|
+
if self.dir_only and not is_dir:
|
|
162
|
+
return False
|
|
163
|
+
if self.prefix:
|
|
164
|
+
if not rel.startswith(self.prefix):
|
|
165
|
+
return False
|
|
166
|
+
rel = rel[len(self.prefix) :]
|
|
167
|
+
if self.anchored:
|
|
168
|
+
return self.regex.match(rel) is not None
|
|
169
|
+
# Patterns without a slash match any name at any depth.
|
|
170
|
+
return self.regex.match(PurePosixPath(rel).name) is not None
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def _translate(pattern: str) -> str:
|
|
174
|
+
"""Turn a gitignore glob into a regex. Good enough for real-world files."""
|
|
175
|
+
out: list[str] = []
|
|
176
|
+
i = 0
|
|
177
|
+
while i < len(pattern):
|
|
178
|
+
char = pattern[i]
|
|
179
|
+
if char == "*" and pattern[i : i + 2] == "**":
|
|
180
|
+
if pattern[i : i + 3] == "**/":
|
|
181
|
+
out.append("(?:.*/)?") # any number of directories
|
|
182
|
+
i += 3
|
|
183
|
+
else:
|
|
184
|
+
out.append(".*")
|
|
185
|
+
i += 2
|
|
186
|
+
elif char == "*":
|
|
187
|
+
out.append("[^/]*")
|
|
188
|
+
i += 1
|
|
189
|
+
elif char == "?":
|
|
190
|
+
out.append("[^/]")
|
|
191
|
+
i += 1
|
|
192
|
+
elif char == "[":
|
|
193
|
+
end = pattern.find("]", i)
|
|
194
|
+
if end == -1:
|
|
195
|
+
out.append(re.escape(char))
|
|
196
|
+
i += 1
|
|
197
|
+
else:
|
|
198
|
+
body = pattern[i + 1 : end]
|
|
199
|
+
if body.startswith("!"): # git uses ! where regex uses ^
|
|
200
|
+
body = "^" + body[1:]
|
|
201
|
+
out.append("[" + body + "]")
|
|
202
|
+
i = end + 1
|
|
203
|
+
else:
|
|
204
|
+
out.append(re.escape(char))
|
|
205
|
+
i += 1
|
|
206
|
+
return "".join(out)
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _compile_rule(line: str, base_rel: str) -> _Rule | None:
|
|
210
|
+
pattern = line.rstrip()
|
|
211
|
+
if not pattern or pattern.startswith("#"):
|
|
212
|
+
return None
|
|
213
|
+
negated = pattern.startswith("!")
|
|
214
|
+
if negated:
|
|
215
|
+
pattern = pattern[1:]
|
|
216
|
+
dir_only = pattern.endswith("/")
|
|
217
|
+
if dir_only:
|
|
218
|
+
pattern = pattern[:-1]
|
|
219
|
+
if not pattern:
|
|
220
|
+
return None
|
|
221
|
+
anchored = pattern.startswith("/") or "/" in pattern
|
|
222
|
+
pattern = pattern.removeprefix("/")
|
|
223
|
+
prefix = "" if base_rel == "." else base_rel + "/"
|
|
224
|
+
return _Rule(
|
|
225
|
+
prefix=prefix,
|
|
226
|
+
regex=re.compile("^" + _translate(pattern) + "$"),
|
|
227
|
+
negated=negated,
|
|
228
|
+
dir_only=dir_only,
|
|
229
|
+
anchored=anchored,
|
|
230
|
+
)
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
def _load_gitignore(directory: Path, base_rel: str) -> list[_Rule]:
|
|
234
|
+
path = directory / ".gitignore"
|
|
235
|
+
if not path.is_file():
|
|
236
|
+
return []
|
|
237
|
+
try:
|
|
238
|
+
if path.stat().st_size > 64 * 1024:
|
|
239
|
+
log.debug("skipping oversized .gitignore: %s", path)
|
|
240
|
+
return []
|
|
241
|
+
text = path.read_text(encoding="utf-8", errors="replace")
|
|
242
|
+
except OSError as exc:
|
|
243
|
+
log.debug("cannot read %s: %s", path, exc)
|
|
244
|
+
return []
|
|
245
|
+
rules = []
|
|
246
|
+
for line in text.splitlines():
|
|
247
|
+
rule = _compile_rule(line, base_rel)
|
|
248
|
+
if rule is not None:
|
|
249
|
+
rules.append(rule)
|
|
250
|
+
return rules
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
def _is_ignored(rel: str, rules: list[_Rule], is_dir: bool) -> bool:
|
|
254
|
+
ignored = False
|
|
255
|
+
for rule in rules: # last matching rule wins, git-style
|
|
256
|
+
if rule.matches(rel, is_dir):
|
|
257
|
+
ignored = not rule.negated
|
|
258
|
+
return ignored
|
|
259
|
+
|
|
260
|
+
|
|
261
|
+
def _ancestor_rules(root: Path, start: Path) -> list[_Rule]:
|
|
262
|
+
"""Gitignore rules from root down to (but not including) start."""
|
|
263
|
+
if start == root:
|
|
264
|
+
return []
|
|
265
|
+
rel = start.relative_to(root)
|
|
266
|
+
rules: list[_Rule] = []
|
|
267
|
+
for i in range(len(rel.parts)):
|
|
268
|
+
base = "/".join(rel.parts[:i]) or "."
|
|
269
|
+
rules.extend(_load_gitignore(root.joinpath(*rel.parts[:i]), base))
|
|
270
|
+
return rules
|
|
271
|
+
|
|
272
|
+
|
|
273
|
+
def _language_of(name: str) -> Language | None:
|
|
274
|
+
return SOURCE_EXTENSIONS.get(PurePosixPath(name).suffix.lower())
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def _walk(
|
|
278
|
+
root: Path,
|
|
279
|
+
current: Path,
|
|
280
|
+
base_rel: str,
|
|
281
|
+
inherited: list[_Rule],
|
|
282
|
+
out: list[FileEntry],
|
|
283
|
+
config: Config,
|
|
284
|
+
) -> None:
|
|
285
|
+
rules = list(inherited) + _load_gitignore(current, base_rel)
|
|
286
|
+
try:
|
|
287
|
+
entries = sorted(current.iterdir())
|
|
288
|
+
except OSError as exc:
|
|
289
|
+
log.debug("cannot list %s: %s", current, exc)
|
|
290
|
+
return
|
|
291
|
+
for entry in entries:
|
|
292
|
+
try:
|
|
293
|
+
rel = entry.relative_to(root).as_posix()
|
|
294
|
+
if entry.is_dir():
|
|
295
|
+
if entry.is_symlink() or entry.name in DENIED_DIRS or is_denied(rel):
|
|
296
|
+
continue
|
|
297
|
+
if _is_ignored(rel, rules, is_dir=True):
|
|
298
|
+
continue
|
|
299
|
+
_walk(root, entry, rel, rules, out, config)
|
|
300
|
+
elif entry.is_file():
|
|
301
|
+
if is_denied(rel) or _is_ignored(rel, rules, is_dir=False):
|
|
302
|
+
continue
|
|
303
|
+
size = entry.stat().st_size
|
|
304
|
+
if size > config.max_file_bytes:
|
|
305
|
+
continue
|
|
306
|
+
language = _language_of(entry.name)
|
|
307
|
+
if language is None and not _is_build_file(entry.name):
|
|
308
|
+
continue
|
|
309
|
+
out.append(FileEntry(path=rel, size=size, language=language))
|
|
310
|
+
except OSError as exc:
|
|
311
|
+
log.debug("skipping %s: %s", entry, exc)
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def _is_build_file(name: str) -> bool:
|
|
315
|
+
return name.lower() in BUILD_FILE_NAMES
|
|
316
|
+
|
|
317
|
+
|
|
318
|
+
def collect_files(
|
|
319
|
+
root: Path, config: Config, *, start: Path | None = None
|
|
320
|
+
) -> list[FileEntry]:
|
|
321
|
+
"""Source and build files under start (default: root), sorted by path.
|
|
322
|
+
|
|
323
|
+
Honours every .gitignore from the root down to start, plus the always-on
|
|
324
|
+
deny rules and the per-file size cap.
|
|
325
|
+
"""
|
|
326
|
+
root = Path(root).resolve()
|
|
327
|
+
start = Path(start).resolve() if start else root
|
|
328
|
+
if not is_within(root, start):
|
|
329
|
+
raise ValueError(f"start is outside the project root: {start}")
|
|
330
|
+
rules = _ancestor_rules(root, start)
|
|
331
|
+
out: list[FileEntry] = []
|
|
332
|
+
base = "." if start == root else start.relative_to(root).as_posix()
|
|
333
|
+
_walk(root, start, base, rules, out, config)
|
|
334
|
+
out.sort(key=lambda entry: entry.path)
|
|
335
|
+
return out
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
def limit_entries(entries: list[FileEntry], config: Config) -> list[FileEntry]:
|
|
339
|
+
"""Cut a file list down to the context budget and the file cap."""
|
|
340
|
+
limited: list[FileEntry] = []
|
|
341
|
+
total = 0
|
|
342
|
+
for entry in entries:
|
|
343
|
+
if len(limited) >= config.max_files or total + entry.size > config.max_context_bytes:
|
|
344
|
+
break
|
|
345
|
+
limited.append(entry)
|
|
346
|
+
total += entry.size
|
|
347
|
+
return limited
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
def _detect_languages(entries: list[FileEntry]) -> list[Language]:
|
|
351
|
+
found = {entry.language for entry in entries if entry.language}
|
|
352
|
+
return [lang for lang in Language if lang in found]
|
|
353
|
+
|
|
354
|
+
|
|
355
|
+
def _detect_build_systems(entries: list[FileEntry]) -> list[BuildSystem]:
|
|
356
|
+
found: set[BuildSystem] = set()
|
|
357
|
+
for entry in entries:
|
|
358
|
+
name = PurePosixPath(entry.path).name.lower()
|
|
359
|
+
if name in MAKE_NAMES:
|
|
360
|
+
found.add(BuildSystem.MAKE)
|
|
361
|
+
elif name in CMAKE_NAMES:
|
|
362
|
+
found.add(BuildSystem.CMAKE)
|
|
363
|
+
elif name in CARGO_NAMES:
|
|
364
|
+
found.add(BuildSystem.CARGO)
|
|
365
|
+
elif name in PYTEST_NAMES:
|
|
366
|
+
found.add(BuildSystem.PYTEST)
|
|
367
|
+
return [build for build in BuildSystem if build in found]
|
|
368
|
+
|
|
369
|
+
|
|
370
|
+
def discover(root: Path, config: Config | None = None) -> ProjectContext:
|
|
371
|
+
config = config or Config()
|
|
372
|
+
root = Path(root).resolve()
|
|
373
|
+
candidates = collect_files(root, config)
|
|
374
|
+
return ProjectContext(
|
|
375
|
+
root=root,
|
|
376
|
+
languages=_detect_languages(candidates),
|
|
377
|
+
build_systems=_detect_build_systems(candidates),
|
|
378
|
+
files=limit_entries(candidates, config),
|
|
379
|
+
)
|
ferrum/model.py
ADDED
|
@@ -0,0 +1,302 @@
|
|
|
1
|
+
"""Model access through an OpenAI-compatible chat endpoint.
|
|
2
|
+
|
|
3
|
+
Ollama exposes one at http://localhost:11434/v1; vLLM and friends do too.
|
|
4
|
+
Only the endpoints actually needed here are implemented: chat completion
|
|
5
|
+
with tool calls, and a best-effort model listing for error hints.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import json
|
|
11
|
+
import logging
|
|
12
|
+
import time
|
|
13
|
+
import urllib.error
|
|
14
|
+
import urllib.request
|
|
15
|
+
from abc import ABC, abstractmethod
|
|
16
|
+
from dataclasses import dataclass, field
|
|
17
|
+
from typing import Any
|
|
18
|
+
|
|
19
|
+
from ferrum import __version__
|
|
20
|
+
|
|
21
|
+
log = logging.getLogger(__name__)
|
|
22
|
+
|
|
23
|
+
MAX_ATTEMPTS = 3
|
|
24
|
+
RETRYABLE_STATUS = {500, 502, 503, 504}
|
|
25
|
+
# Free tiers meter per minute; waiting out the window beats failing the run.
|
|
26
|
+
RATE_LIMIT_BACKOFF = (15, 30)
|
|
27
|
+
|
|
28
|
+
MALFORMED_JSON = "malformed_json"
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class ProviderError(Exception):
|
|
32
|
+
"""The model could not be reached or answered unusably."""
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class AuthenticationError(ProviderError):
|
|
36
|
+
pass
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class RateLimitError(ProviderError):
|
|
40
|
+
pass
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
@dataclass(frozen=True)
|
|
44
|
+
class ToolCall:
|
|
45
|
+
id: str
|
|
46
|
+
name: str
|
|
47
|
+
arguments: dict[str, Any] = field(default_factory=dict)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
@dataclass(frozen=True)
|
|
51
|
+
class ModelResponse:
|
|
52
|
+
content: str
|
|
53
|
+
tool_calls: list[ToolCall] = field(default_factory=list)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
class ModelProvider(ABC):
|
|
57
|
+
@abstractmethod
|
|
58
|
+
def complete(
|
|
59
|
+
self, messages: list[dict[str, Any]], tools: list[dict[str, Any]]
|
|
60
|
+
) -> ModelResponse:
|
|
61
|
+
"""One round trip: messages in, content and tool calls out."""
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
class OpenAICompatibleProvider(ModelProvider):
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
base_url: str,
|
|
68
|
+
model: str,
|
|
69
|
+
api_key: str = "",
|
|
70
|
+
timeout: int = 60,
|
|
71
|
+
max_tokens: int = 2048,
|
|
72
|
+
) -> None:
|
|
73
|
+
self.base_url = base_url.rstrip("/")
|
|
74
|
+
self.model = model
|
|
75
|
+
self.api_key = api_key
|
|
76
|
+
self.timeout = timeout
|
|
77
|
+
self.max_tokens = max_tokens
|
|
78
|
+
|
|
79
|
+
def complete(
|
|
80
|
+
self, messages: list[dict[str, Any]], tools: list[dict[str, Any]]
|
|
81
|
+
) -> ModelResponse:
|
|
82
|
+
payload: dict[str, Any] = {
|
|
83
|
+
"model": self.model,
|
|
84
|
+
"messages": messages,
|
|
85
|
+
"stream": False,
|
|
86
|
+
# Without an explicit cap some servers truncate mid-tool-call.
|
|
87
|
+
"max_tokens": self.max_tokens,
|
|
88
|
+
"temperature": 0.2,
|
|
89
|
+
}
|
|
90
|
+
if tools:
|
|
91
|
+
payload["tools"] = tools
|
|
92
|
+
body = json.dumps(payload).encode("utf-8")
|
|
93
|
+
headers = {
|
|
94
|
+
"Content-Type": "application/json",
|
|
95
|
+
# Some edges (Cloudflare) reject urllib's default signature.
|
|
96
|
+
"User-Agent": f"ferrum/{__version__}",
|
|
97
|
+
}
|
|
98
|
+
if self.api_key:
|
|
99
|
+
headers["Authorization"] = f"Bearer {self.api_key}"
|
|
100
|
+
request = urllib.request.Request(
|
|
101
|
+
f"{self.base_url}/chat/completions",
|
|
102
|
+
data=body,
|
|
103
|
+
headers=headers,
|
|
104
|
+
method="POST",
|
|
105
|
+
)
|
|
106
|
+
raw = self._send(request)
|
|
107
|
+
return self._parse(raw)
|
|
108
|
+
|
|
109
|
+
def _send(self, request: urllib.request.Request) -> bytes:
|
|
110
|
+
# Gateways flake; retry transient failures before giving up.
|
|
111
|
+
for attempt in range(1, MAX_ATTEMPTS + 1):
|
|
112
|
+
try:
|
|
113
|
+
with urllib.request.urlopen(request, timeout=self.timeout) as response:
|
|
114
|
+
return response.read()
|
|
115
|
+
except urllib.error.HTTPError as exc:
|
|
116
|
+
if exc.code == 429 and attempt < MAX_ATTEMPTS:
|
|
117
|
+
wait = RATE_LIMIT_BACKOFF[attempt - 1]
|
|
118
|
+
log.warning(
|
|
119
|
+
"model server HTTP 429 (rate limited),"
|
|
120
|
+
" waiting %ds (attempt %d/%d)",
|
|
121
|
+
wait, attempt, MAX_ATTEMPTS,
|
|
122
|
+
)
|
|
123
|
+
exc.close()
|
|
124
|
+
time.sleep(wait)
|
|
125
|
+
continue
|
|
126
|
+
if exc.code in RETRYABLE_STATUS and attempt < MAX_ATTEMPTS:
|
|
127
|
+
log.warning(
|
|
128
|
+
"model server HTTP %s (attempt %d/%d), retrying",
|
|
129
|
+
exc.code, attempt, MAX_ATTEMPTS,
|
|
130
|
+
)
|
|
131
|
+
exc.close()
|
|
132
|
+
time.sleep(attempt)
|
|
133
|
+
continue
|
|
134
|
+
raise self._http_error(exc) from exc
|
|
135
|
+
except TimeoutError as exc:
|
|
136
|
+
raise ProviderError(
|
|
137
|
+
f"model at {self.base_url} did not answer within "
|
|
138
|
+
f"{self.timeout}s"
|
|
139
|
+
) from exc
|
|
140
|
+
except urllib.error.URLError as exc:
|
|
141
|
+
if attempt < MAX_ATTEMPTS:
|
|
142
|
+
log.warning(
|
|
143
|
+
"cannot reach model (attempt %d/%d), retrying: %s",
|
|
144
|
+
attempt, MAX_ATTEMPTS, exc.reason,
|
|
145
|
+
)
|
|
146
|
+
time.sleep(attempt)
|
|
147
|
+
continue
|
|
148
|
+
raise ProviderError(
|
|
149
|
+
f"cannot reach model at {self.base_url} "
|
|
150
|
+
f"(is the server running?): {exc.reason}"
|
|
151
|
+
) from exc
|
|
152
|
+
except OSError as exc:
|
|
153
|
+
raise ProviderError(
|
|
154
|
+
f"cannot reach model at {self.base_url}: {exc}"
|
|
155
|
+
) from exc
|
|
156
|
+
raise ProviderError(f"cannot reach model at {self.base_url}") # pragma: no cover
|
|
157
|
+
|
|
158
|
+
def _http_error(self, exc: urllib.error.HTTPError) -> ProviderError:
|
|
159
|
+
try:
|
|
160
|
+
detail = exc.read().decode("utf-8", "replace")[:300]
|
|
161
|
+
except OSError:
|
|
162
|
+
detail = ""
|
|
163
|
+
if exc.code in (401, 403):
|
|
164
|
+
return AuthenticationError(
|
|
165
|
+
f"the model server refused the request (HTTP {exc.code}): {detail}"
|
|
166
|
+
)
|
|
167
|
+
if exc.code == 429:
|
|
168
|
+
return RateLimitError(
|
|
169
|
+
f"rate limited by the model server (HTTP 429): {detail}"
|
|
170
|
+
)
|
|
171
|
+
return ProviderError(f"model server returned HTTP {exc.code}: {detail}")
|
|
172
|
+
|
|
173
|
+
def _parse(self, raw: bytes) -> ModelResponse:
|
|
174
|
+
try:
|
|
175
|
+
message = json.loads(raw)["choices"][0]["message"]
|
|
176
|
+
except (ValueError, KeyError, IndexError, TypeError) as exc:
|
|
177
|
+
raise ProviderError(f"malformed response from model: {exc}") from exc
|
|
178
|
+
content = message.get("content") or ""
|
|
179
|
+
if not isinstance(content, str):
|
|
180
|
+
content = str(content)
|
|
181
|
+
calls: list[ToolCall] = []
|
|
182
|
+
for entry in message.get("tool_calls") or []:
|
|
183
|
+
if not isinstance(entry, dict):
|
|
184
|
+
continue
|
|
185
|
+
function = entry.get("function") or {}
|
|
186
|
+
name = function.get("name") or ""
|
|
187
|
+
raw_args = function.get("arguments") or "{}"
|
|
188
|
+
if isinstance(raw_args, dict):
|
|
189
|
+
arguments: dict[str, Any] = raw_args
|
|
190
|
+
else:
|
|
191
|
+
try:
|
|
192
|
+
parsed = json.loads(raw_args)
|
|
193
|
+
arguments = parsed if isinstance(parsed, dict) else {"_": parsed}
|
|
194
|
+
except ValueError:
|
|
195
|
+
# Small models emit unescaped newlines; let the agent
|
|
196
|
+
# feed the error back instead of killing the session.
|
|
197
|
+
arguments = {MALFORMED_JSON: str(raw_args)}
|
|
198
|
+
arguments = _normalize_arguments(arguments)
|
|
199
|
+
calls.append(
|
|
200
|
+
ToolCall(
|
|
201
|
+
id=str(entry.get("id") or name or len(calls)),
|
|
202
|
+
name=str(name),
|
|
203
|
+
arguments=arguments,
|
|
204
|
+
)
|
|
205
|
+
)
|
|
206
|
+
return ModelResponse(content=content, tool_calls=calls)
|
|
207
|
+
|
|
208
|
+
|
|
209
|
+
def _normalize_arguments(args: dict[str, Any]) -> dict[str, Any]:
|
|
210
|
+
"""Small models often echo the tools schema instead of real values.
|
|
211
|
+
|
|
212
|
+
Unwrap the shapes they produce; anything schema-shaped becomes a
|
|
213
|
+
MALFORMED marker so the harness can ask the model to retry with values.
|
|
214
|
+
"""
|
|
215
|
+
single = args.get("arguments")
|
|
216
|
+
if set(args) == {"arguments"} and isinstance(single, dict):
|
|
217
|
+
return _normalize_arguments(single)
|
|
218
|
+
parameters = args.get("parameters")
|
|
219
|
+
if isinstance(parameters, dict):
|
|
220
|
+
if "properties" in parameters:
|
|
221
|
+
return {MALFORMED_JSON: json.dumps(args)}
|
|
222
|
+
return parameters
|
|
223
|
+
if "properties" in args or "required" in args:
|
|
224
|
+
return {MALFORMED_JSON: json.dumps(args)}
|
|
225
|
+
return args
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def extract_tool_calls(content: str) -> list[ToolCall]:
|
|
229
|
+
"""Fallback for models that write tool calls as plain JSON in their text.
|
|
230
|
+
|
|
231
|
+
Scans for JSON objects that name a tool; anything else in the prose is
|
|
232
|
+
ignored. The harness reports unknown names back, so invented tools get
|
|
233
|
+
corrected on the next turn.
|
|
234
|
+
"""
|
|
235
|
+
if not content:
|
|
236
|
+
return []
|
|
237
|
+
decoder = json.JSONDecoder()
|
|
238
|
+
calls: list[ToolCall] = []
|
|
239
|
+
position = 0
|
|
240
|
+
while True:
|
|
241
|
+
start = content.find("{", position)
|
|
242
|
+
if start == -1:
|
|
243
|
+
break
|
|
244
|
+
try:
|
|
245
|
+
obj, end = decoder.raw_decode(content, start)
|
|
246
|
+
except ValueError:
|
|
247
|
+
position = start + 1
|
|
248
|
+
continue
|
|
249
|
+
position = end
|
|
250
|
+
if not isinstance(obj, dict):
|
|
251
|
+
continue
|
|
252
|
+
name = obj.get("name") or obj.get("tool")
|
|
253
|
+
if not name and isinstance(obj.get("function"), str):
|
|
254
|
+
name = obj["function"]
|
|
255
|
+
if not isinstance(name, str) or not name:
|
|
256
|
+
continue
|
|
257
|
+
arguments = obj.get("arguments", obj.get("parameters", obj.get("args", {})))
|
|
258
|
+
if isinstance(arguments, str):
|
|
259
|
+
try:
|
|
260
|
+
parsed = json.loads(arguments)
|
|
261
|
+
except ValueError:
|
|
262
|
+
arguments = {MALFORMED_JSON: arguments}
|
|
263
|
+
else:
|
|
264
|
+
arguments = parsed if isinstance(parsed, dict) else {"_": parsed}
|
|
265
|
+
elif not isinstance(arguments, dict):
|
|
266
|
+
arguments = {"_": arguments}
|
|
267
|
+
arguments = _normalize_arguments(arguments)
|
|
268
|
+
calls.append(ToolCall(id=f"text-{len(calls)}", name=name, arguments=arguments))
|
|
269
|
+
return calls
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def probe_endpoint(
|
|
273
|
+
base_url: str, timeout: int = 5, api_key: str = ""
|
|
274
|
+
) -> tuple[list[str], str | None]:
|
|
275
|
+
"""(models, error) for GET /models. Unlike list_models, failures are
|
|
276
|
+
reported instead of swallowed, so `ferrum doctor` can show why."""
|
|
277
|
+
url = base_url.rstrip("/") + "/models"
|
|
278
|
+
request = urllib.request.Request(url)
|
|
279
|
+
request.add_header("User-Agent", f"ferrum/{__version__}")
|
|
280
|
+
if api_key:
|
|
281
|
+
request.add_header("Authorization", f"Bearer {api_key}")
|
|
282
|
+
try:
|
|
283
|
+
with urllib.request.urlopen(request, timeout=timeout) as response:
|
|
284
|
+
data = json.loads(response.read())
|
|
285
|
+
models = [str(m["id"]) for m in data.get("data", []) if m.get("id")]
|
|
286
|
+
except urllib.error.HTTPError as exc:
|
|
287
|
+
exc.close()
|
|
288
|
+
return [], f"HTTP {exc.code}"
|
|
289
|
+
except Exception as exc: # noqa: BLE001 - report any failure, raise none
|
|
290
|
+
log.debug("model listing failed: %s", exc)
|
|
291
|
+
return [], str(exc) or type(exc).__name__
|
|
292
|
+
return models, None
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def list_models(base_url: str, timeout: int = 5, api_key: str = "") -> list[str]:
|
|
296
|
+
"""Best-effort model listing for hints; never raises.
|
|
297
|
+
|
|
298
|
+
Auth is optional: some endpoints (Cloudflare-fronted clouds) refuse
|
|
299
|
+
anonymous GET /models, while local Ollama ignores the header.
|
|
300
|
+
"""
|
|
301
|
+
models, _error = probe_endpoint(base_url, timeout=timeout, api_key=api_key)
|
|
302
|
+
return models
|