driangle-agentrunner 0.0.1__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.
- agentrunner/__init__.py +35 -0
- agentrunner/claudecode/__init__.py +11 -0
- agentrunner/claudecode/args.py +51 -0
- agentrunner/claudecode/mapping.py +59 -0
- agentrunner/claudecode/options.py +33 -0
- agentrunner/claudecode/parser.py +152 -0
- agentrunner/claudecode/process.py +68 -0
- agentrunner/claudecode/runner.py +270 -0
- agentrunner/claudecode/types.py +131 -0
- agentrunner/claudecode/version.py +47 -0
- agentrunner/errors.py +36 -0
- agentrunner/ollama/__init__.py +14 -0
- agentrunner/ollama/accessors.py +22 -0
- agentrunner/ollama/options.py +41 -0
- agentrunner/ollama/runner.py +327 -0
- agentrunner/ollama/types.py +87 -0
- agentrunner/types.py +188 -0
- driangle_agentrunner-0.0.1.dist-info/METADATA +219 -0
- driangle_agentrunner-0.0.1.dist-info/RECORD +20 -0
- driangle_agentrunner-0.0.1.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
"""Types for Claude Code CLI stream-json output."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class ContentBlock:
|
|
11
|
+
"""One block inside an assistant message."""
|
|
12
|
+
|
|
13
|
+
type: str = ""
|
|
14
|
+
text: str | None = None
|
|
15
|
+
thinking: str | None = None
|
|
16
|
+
name: str | None = None
|
|
17
|
+
input: Any = None
|
|
18
|
+
content: Any = None
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass
|
|
22
|
+
class StreamUsage:
|
|
23
|
+
"""Token counts from streaming events."""
|
|
24
|
+
|
|
25
|
+
input_tokens: int = 0
|
|
26
|
+
output_tokens: int = 0
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass
|
|
30
|
+
class ResultUsage:
|
|
31
|
+
"""Token counts from the final result message (includes cache fields)."""
|
|
32
|
+
|
|
33
|
+
input_tokens: int = 0
|
|
34
|
+
output_tokens: int = 0
|
|
35
|
+
cache_creation_input_tokens: int = 0
|
|
36
|
+
cache_read_input_tokens: int = 0
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass
|
|
40
|
+
class Delta:
|
|
41
|
+
"""Incremental data in delta events."""
|
|
42
|
+
|
|
43
|
+
type: str | None = None
|
|
44
|
+
text: str | None = None
|
|
45
|
+
thinking: str | None = None
|
|
46
|
+
partial_json: str | None = None
|
|
47
|
+
stop_reason: str | None = None
|
|
48
|
+
stop_sequence: str | None = None
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@dataclass
|
|
52
|
+
class ContentBlockInfo:
|
|
53
|
+
"""Content block info in content_block_start events."""
|
|
54
|
+
|
|
55
|
+
type: str = ""
|
|
56
|
+
name: str | None = None
|
|
57
|
+
id: str | None = None
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
@dataclass
|
|
61
|
+
class MessageStartData:
|
|
62
|
+
"""Message metadata from a message_start event."""
|
|
63
|
+
|
|
64
|
+
model: str = ""
|
|
65
|
+
id: str = ""
|
|
66
|
+
usage: StreamUsage | None = None
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
@dataclass
|
|
70
|
+
class StreamEventInner:
|
|
71
|
+
"""Parsed inner event from a stream_event line."""
|
|
72
|
+
|
|
73
|
+
type: str = ""
|
|
74
|
+
message: MessageStartData | None = None
|
|
75
|
+
index: int | None = None
|
|
76
|
+
content_block: ContentBlockInfo | None = None
|
|
77
|
+
delta: Delta | None = None
|
|
78
|
+
usage: StreamUsage | None = None
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
@dataclass
|
|
82
|
+
class RateLimitInfo:
|
|
83
|
+
"""Rate limit details from rate_limit_event messages."""
|
|
84
|
+
|
|
85
|
+
status: str = ""
|
|
86
|
+
rate_limit_type: str | None = None
|
|
87
|
+
utilization: float | None = None
|
|
88
|
+
resets_at: float | None = None
|
|
89
|
+
is_using_overage: bool | None = None
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
@dataclass
|
|
93
|
+
class AssistantMessage:
|
|
94
|
+
"""Nested 'message' object inside assistant-type stream lines."""
|
|
95
|
+
|
|
96
|
+
model: str | None = None
|
|
97
|
+
id: str | None = None
|
|
98
|
+
content: list[ContentBlock] = field(default_factory=list)
|
|
99
|
+
stop_reason: str | None = None
|
|
100
|
+
usage: StreamUsage | None = None
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
@dataclass
|
|
104
|
+
class StreamMessage:
|
|
105
|
+
"""Top-level envelope for all Claude stream-json lines."""
|
|
106
|
+
|
|
107
|
+
type: str = ""
|
|
108
|
+
subtype: str | None = None
|
|
109
|
+
content: list[ContentBlock] = field(default_factory=list)
|
|
110
|
+
message: AssistantMessage | None = None
|
|
111
|
+
|
|
112
|
+
# Result fields.
|
|
113
|
+
result: str | None = None
|
|
114
|
+
is_error: bool | None = None
|
|
115
|
+
total_cost_usd: float | None = None
|
|
116
|
+
duration_ms: float | None = None
|
|
117
|
+
duration_api_ms: float | None = None
|
|
118
|
+
num_turns: int | None = None
|
|
119
|
+
session_id: str | None = None
|
|
120
|
+
model: str | None = None
|
|
121
|
+
usage: ResultUsage | None = None
|
|
122
|
+
|
|
123
|
+
# System/init fields.
|
|
124
|
+
tools: list[Any] | None = None
|
|
125
|
+
|
|
126
|
+
# Rate limit info.
|
|
127
|
+
rate_limit_info: RateLimitInfo | None = None
|
|
128
|
+
|
|
129
|
+
# Stream event fields.
|
|
130
|
+
event: StreamEventInner | None = None
|
|
131
|
+
parent_tool_use_id: str | None = None
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
"""CLI version detection and compatibility check."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import re
|
|
7
|
+
|
|
8
|
+
from ..errors import NotFoundError
|
|
9
|
+
|
|
10
|
+
# Supported Claude Code CLI version range.
|
|
11
|
+
MIN_VERSION = "1.0.12"
|
|
12
|
+
|
|
13
|
+
_VERSION_RE = re.compile(r"(\d+\.\d+\.\d+)")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _parse_version(version_str: str) -> tuple[int, ...]:
|
|
17
|
+
"""Parse a semver string into a comparable tuple."""
|
|
18
|
+
return tuple(int(p) for p in version_str.split("."))
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
async def check_version(binary: str) -> str:
|
|
22
|
+
"""Run ``<binary> --version`` and verify it meets the minimum requirement.
|
|
23
|
+
|
|
24
|
+
Returns the detected version string.
|
|
25
|
+
Raises ``NotFoundError`` if the binary is missing or the version is too old.
|
|
26
|
+
"""
|
|
27
|
+
try:
|
|
28
|
+
proc = await asyncio.create_subprocess_exec(
|
|
29
|
+
binary,
|
|
30
|
+
"--version",
|
|
31
|
+
stdout=asyncio.subprocess.PIPE,
|
|
32
|
+
stderr=asyncio.subprocess.PIPE,
|
|
33
|
+
)
|
|
34
|
+
stdout, _ = await proc.communicate()
|
|
35
|
+
except FileNotFoundError:
|
|
36
|
+
raise NotFoundError(f"{binary}: command not found")
|
|
37
|
+
|
|
38
|
+
output = stdout.decode("utf-8", errors="replace").strip()
|
|
39
|
+
match = _VERSION_RE.search(output)
|
|
40
|
+
if not match:
|
|
41
|
+
raise NotFoundError(f"could not parse version from `{binary} --version`: {output!r}")
|
|
42
|
+
|
|
43
|
+
version = match.group(1)
|
|
44
|
+
if _parse_version(version) < _parse_version(MIN_VERSION):
|
|
45
|
+
raise NotFoundError(f"{binary} version {version} is below minimum supported {MIN_VERSION}")
|
|
46
|
+
|
|
47
|
+
return version
|
agentrunner/errors.py
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
"""Exception hierarchy for runner errors."""
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
class RunnerError(Exception):
|
|
5
|
+
"""Base class for all runner errors."""
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class NotFoundError(RunnerError):
|
|
9
|
+
"""Runner binary or API endpoint is not reachable."""
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class TimeoutError(RunnerError):
|
|
13
|
+
"""Execution exceeded the configured timeout."""
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class NonZeroExitError(RunnerError):
|
|
17
|
+
"""CLI process exited with a non-zero code."""
|
|
18
|
+
|
|
19
|
+
def __init__(self, exit_code: int, message: str) -> None:
|
|
20
|
+
super().__init__(message)
|
|
21
|
+
self.exit_code = exit_code
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class ParseError(RunnerError):
|
|
25
|
+
"""Failed to parse runner output."""
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class CancelledError(RunnerError):
|
|
29
|
+
"""Execution was cancelled by the caller."""
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class NoResultError(RunnerError):
|
|
33
|
+
"""Stream ended without a result message."""
|
|
34
|
+
|
|
35
|
+
def __init__(self) -> None:
|
|
36
|
+
super().__init__("no result in output")
|
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
"""Ollama runner — talks to the Ollama HTTP API for local model inference."""
|
|
2
|
+
|
|
3
|
+
from .accessors import message_text, message_thinking
|
|
4
|
+
from .options import OllamaRunnerConfig, OllamaRunOptions
|
|
5
|
+
from .runner import OllamaRunner, OllamaSession
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"OllamaRunner",
|
|
9
|
+
"OllamaRunOptions",
|
|
10
|
+
"OllamaRunnerConfig",
|
|
11
|
+
"OllamaSession",
|
|
12
|
+
"message_text",
|
|
13
|
+
"message_thinking",
|
|
14
|
+
]
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
"""Typed accessor functions for Ollama messages."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from ..types import Message
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def message_text(msg: Message) -> str | None:
|
|
9
|
+
"""Return the text content from an assistant or result message, or None."""
|
|
10
|
+
d = msg._raw_dict()
|
|
11
|
+
message = d.get("message", {})
|
|
12
|
+
content = message.get("content")
|
|
13
|
+
if content:
|
|
14
|
+
return content
|
|
15
|
+
return None
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def message_thinking(msg: Message) -> str | None:
|
|
19
|
+
"""Return the thinking content from a message, or None."""
|
|
20
|
+
d = msg._raw_dict()
|
|
21
|
+
message = d.get("message", {})
|
|
22
|
+
return message.get("thinking") or None
|
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
"""Ollama runner configuration and run options."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
from typing import Any, Protocol, runtime_checkable
|
|
7
|
+
|
|
8
|
+
from ..types import RunOptions
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@runtime_checkable
|
|
12
|
+
class Logger(Protocol):
|
|
13
|
+
"""Minimal logger interface (matches stdlib logging.Logger)."""
|
|
14
|
+
|
|
15
|
+
def debug(self, msg: str, *args: Any, **kwargs: Any) -> None: ...
|
|
16
|
+
def error(self, msg: str, *args: Any, **kwargs: Any) -> None: ...
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass
|
|
20
|
+
class OllamaRunnerConfig:
|
|
21
|
+
"""Configuration for creating an Ollama runner."""
|
|
22
|
+
|
|
23
|
+
base_url: str = "http://localhost:11434"
|
|
24
|
+
logger: Logger | None = None
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass
|
|
28
|
+
class OllamaRunOptions(RunOptions):
|
|
29
|
+
"""Ollama-specific options extending common RunOptions."""
|
|
30
|
+
|
|
31
|
+
temperature: float | None = None
|
|
32
|
+
num_ctx: int | None = None
|
|
33
|
+
num_predict: int | None = None
|
|
34
|
+
seed: int | None = None
|
|
35
|
+
stop: list[str] | None = field(default=None)
|
|
36
|
+
top_k: int | None = None
|
|
37
|
+
top_p: float | None = None
|
|
38
|
+
min_p: float | None = None
|
|
39
|
+
format: str | None = None
|
|
40
|
+
keep_alive: str | None = None
|
|
41
|
+
think: bool | None = None
|
|
@@ -0,0 +1,327 @@
|
|
|
1
|
+
"""Ollama runner implementation using the Ollama HTTP API."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import json
|
|
7
|
+
from collections.abc import AsyncIterator
|
|
8
|
+
from http import HTTPStatus
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from ..errors import (
|
|
12
|
+
CancelledError,
|
|
13
|
+
NoResultError,
|
|
14
|
+
NotFoundError,
|
|
15
|
+
ParseError,
|
|
16
|
+
RunnerError,
|
|
17
|
+
TimeoutError,
|
|
18
|
+
)
|
|
19
|
+
from ..types import Message, Result, Usage
|
|
20
|
+
from .options import OllamaRunnerConfig, OllamaRunOptions
|
|
21
|
+
from .types import ChatResponse, ModelOptions
|
|
22
|
+
|
|
23
|
+
DEFAULT_BASE_URL = "http://localhost:11434"
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class OllamaSession:
|
|
27
|
+
"""Session encapsulates a running Ollama API request.
|
|
28
|
+
|
|
29
|
+
Supports ``async for msg in session`` to iterate messages,
|
|
30
|
+
and ``await session.result`` to get the final result.
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
def __init__(
|
|
34
|
+
self,
|
|
35
|
+
config: OllamaRunnerConfig,
|
|
36
|
+
prompt: str,
|
|
37
|
+
options: OllamaRunOptions,
|
|
38
|
+
) -> None:
|
|
39
|
+
self._config = config
|
|
40
|
+
self._prompt = prompt
|
|
41
|
+
self._options = options
|
|
42
|
+
|
|
43
|
+
self._loop = asyncio.get_running_loop()
|
|
44
|
+
self._queue: asyncio.Queue[Message | None] = asyncio.Queue()
|
|
45
|
+
self._result_future: asyncio.Future[Result] = self._loop.create_future()
|
|
46
|
+
self._aborted = False
|
|
47
|
+
self._task: asyncio.Task[None] = asyncio.ensure_future(self._run_request())
|
|
48
|
+
|
|
49
|
+
async def _run_request(self) -> None:
|
|
50
|
+
try:
|
|
51
|
+
base_url = self._config.base_url or DEFAULT_BASE_URL
|
|
52
|
+
body = _build_request_body(self._prompt, self._options)
|
|
53
|
+
|
|
54
|
+
if self._config.logger:
|
|
55
|
+
self._config.logger.debug(
|
|
56
|
+
"executing Ollama API request",
|
|
57
|
+
extra={"method": "POST", "url": f"{base_url}/api/chat"},
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
reader, writer = await _http_post_stream(
|
|
61
|
+
base_url, "/api/chat", json.dumps(body)
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
text_parts: list[str] = []
|
|
65
|
+
final_resp: ChatResponse | None = None
|
|
66
|
+
|
|
67
|
+
try:
|
|
68
|
+
while True:
|
|
69
|
+
if self._aborted:
|
|
70
|
+
break
|
|
71
|
+
|
|
72
|
+
raw_line = await reader.readline()
|
|
73
|
+
if not raw_line:
|
|
74
|
+
break
|
|
75
|
+
|
|
76
|
+
line = raw_line.decode("utf-8", errors="replace").strip()
|
|
77
|
+
if not line:
|
|
78
|
+
continue
|
|
79
|
+
|
|
80
|
+
try:
|
|
81
|
+
chunk = ChatResponse.from_dict(json.loads(line))
|
|
82
|
+
except (json.JSONDecodeError, KeyError) as exc:
|
|
83
|
+
raise ParseError(f"invalid JSON: {line}") from exc
|
|
84
|
+
|
|
85
|
+
if chunk.message.content:
|
|
86
|
+
text_parts.append(chunk.message.content)
|
|
87
|
+
|
|
88
|
+
if chunk.done:
|
|
89
|
+
final_resp = chunk
|
|
90
|
+
|
|
91
|
+
msg = Message(
|
|
92
|
+
type="result" if chunk.done else "assistant",
|
|
93
|
+
raw=line,
|
|
94
|
+
)
|
|
95
|
+
await self._queue.put(msg)
|
|
96
|
+
finally:
|
|
97
|
+
writer.close()
|
|
98
|
+
try:
|
|
99
|
+
await writer.wait_closed()
|
|
100
|
+
except Exception:
|
|
101
|
+
pass
|
|
102
|
+
|
|
103
|
+
if self._aborted:
|
|
104
|
+
self._result_future.set_exception(CancelledError("execution cancelled"))
|
|
105
|
+
return
|
|
106
|
+
|
|
107
|
+
if not final_resp:
|
|
108
|
+
self._result_future.set_exception(NoResultError())
|
|
109
|
+
return
|
|
110
|
+
|
|
111
|
+
usage = Usage(
|
|
112
|
+
input_tokens=final_resp.prompt_eval_count or 0,
|
|
113
|
+
output_tokens=final_resp.eval_count or 0,
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
self._result_future.set_result(
|
|
117
|
+
Result(
|
|
118
|
+
text="".join(text_parts),
|
|
119
|
+
is_error=False,
|
|
120
|
+
exit_code=0,
|
|
121
|
+
usage=usage,
|
|
122
|
+
cost_usd=0.0,
|
|
123
|
+
duration_ms=(
|
|
124
|
+
final_resp.total_duration / 1e6
|
|
125
|
+
if final_resp.total_duration
|
|
126
|
+
else 0.0
|
|
127
|
+
),
|
|
128
|
+
session_id="",
|
|
129
|
+
)
|
|
130
|
+
)
|
|
131
|
+
except (CancelledError, TimeoutError, NotFoundError, ParseError, NoResultError):
|
|
132
|
+
if not self._result_future.done():
|
|
133
|
+
self._result_future.set_exception(
|
|
134
|
+
__import__("sys").exc_info()[1] # type: ignore[arg-type]
|
|
135
|
+
)
|
|
136
|
+
except OSError as exc:
|
|
137
|
+
err = NotFoundError(f"connection failed: {exc}")
|
|
138
|
+
if not self._result_future.done():
|
|
139
|
+
self._result_future.set_exception(err)
|
|
140
|
+
except Exception as exc:
|
|
141
|
+
if not self._result_future.done():
|
|
142
|
+
self._result_future.set_exception(RunnerError(str(exc)))
|
|
143
|
+
finally:
|
|
144
|
+
await self._queue.put(None)
|
|
145
|
+
|
|
146
|
+
def __aiter__(self) -> AsyncIterator[Message]:
|
|
147
|
+
return self._message_iter()
|
|
148
|
+
|
|
149
|
+
async def _message_iter(self) -> AsyncIterator[Message]:
|
|
150
|
+
while True:
|
|
151
|
+
msg = await self._queue.get()
|
|
152
|
+
if msg is None:
|
|
153
|
+
break
|
|
154
|
+
yield msg
|
|
155
|
+
|
|
156
|
+
@property
|
|
157
|
+
def result(self) -> asyncio.Future[Result]:
|
|
158
|
+
return self._result_future
|
|
159
|
+
|
|
160
|
+
def abort(self) -> None:
|
|
161
|
+
self._aborted = True
|
|
162
|
+
|
|
163
|
+
def send(self, input: Any) -> None:
|
|
164
|
+
raise NotImplementedError("send is not supported for Ollama runner")
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
class OllamaRunner:
|
|
168
|
+
"""Ollama runner — talks to the Ollama HTTP API.
|
|
169
|
+
|
|
170
|
+
Construct directly::
|
|
171
|
+
|
|
172
|
+
runner = OllamaRunner()
|
|
173
|
+
runner = OllamaRunner(config=OllamaRunnerConfig(base_url="http://...", logger=logger))
|
|
174
|
+
"""
|
|
175
|
+
|
|
176
|
+
def __init__(self, config: OllamaRunnerConfig | None = None) -> None:
|
|
177
|
+
self._config = config or OllamaRunnerConfig()
|
|
178
|
+
|
|
179
|
+
def start(
|
|
180
|
+
self,
|
|
181
|
+
prompt: str,
|
|
182
|
+
options: OllamaRunOptions | None = None,
|
|
183
|
+
) -> OllamaSession:
|
|
184
|
+
opts = options or OllamaRunOptions()
|
|
185
|
+
if not opts.model:
|
|
186
|
+
raise RunnerError("model is required for Ollama runner")
|
|
187
|
+
|
|
188
|
+
timeout = opts.timeout
|
|
189
|
+
session = OllamaSession(self._config, prompt, opts)
|
|
190
|
+
|
|
191
|
+
if timeout is not None and timeout > 0:
|
|
192
|
+
loop = asyncio.get_running_loop()
|
|
193
|
+
loop.call_later(timeout, session.abort)
|
|
194
|
+
|
|
195
|
+
return session
|
|
196
|
+
|
|
197
|
+
async def run(
|
|
198
|
+
self,
|
|
199
|
+
prompt: str,
|
|
200
|
+
options: OllamaRunOptions | None = None,
|
|
201
|
+
) -> Result:
|
|
202
|
+
session = self.start(prompt, options)
|
|
203
|
+
async for _msg in session:
|
|
204
|
+
pass
|
|
205
|
+
return await session.result
|
|
206
|
+
|
|
207
|
+
async def run_stream(
|
|
208
|
+
self,
|
|
209
|
+
prompt: str,
|
|
210
|
+
options: OllamaRunOptions | None = None,
|
|
211
|
+
) -> OllamaSession:
|
|
212
|
+
return self.start(prompt, options)
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
async def _http_post_stream(
|
|
216
|
+
base_url: str, path: str, body: str
|
|
217
|
+
) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]:
|
|
218
|
+
"""Open a raw HTTP POST connection and return the response body stream.
|
|
219
|
+
|
|
220
|
+
Uses asyncio streams directly to avoid external dependencies.
|
|
221
|
+
Raises NotFoundError on connection failure or HTTP 404.
|
|
222
|
+
"""
|
|
223
|
+
from urllib.parse import urlparse
|
|
224
|
+
|
|
225
|
+
parsed = urlparse(base_url)
|
|
226
|
+
host = parsed.hostname or "localhost"
|
|
227
|
+
port = parsed.port or 80
|
|
228
|
+
use_ssl = parsed.scheme == "https"
|
|
229
|
+
|
|
230
|
+
try:
|
|
231
|
+
if use_ssl:
|
|
232
|
+
import ssl
|
|
233
|
+
|
|
234
|
+
ctx = ssl.create_default_context()
|
|
235
|
+
reader, writer = await asyncio.open_connection(host, port, ssl=ctx)
|
|
236
|
+
else:
|
|
237
|
+
reader, writer = await asyncio.open_connection(host, port)
|
|
238
|
+
except OSError as exc:
|
|
239
|
+
raise NotFoundError(f"connection failed: {exc}") from exc
|
|
240
|
+
|
|
241
|
+
# Send HTTP request.
|
|
242
|
+
body_bytes = body.encode("utf-8")
|
|
243
|
+
request_lines = (
|
|
244
|
+
f"POST {path} HTTP/1.1\r\n"
|
|
245
|
+
f"Host: {host}:{port}\r\n"
|
|
246
|
+
f"Content-Type: application/json\r\n"
|
|
247
|
+
f"Content-Length: {len(body_bytes)}\r\n"
|
|
248
|
+
f"Connection: close\r\n"
|
|
249
|
+
f"\r\n"
|
|
250
|
+
)
|
|
251
|
+
writer.write(request_lines.encode("utf-8"))
|
|
252
|
+
writer.write(body_bytes)
|
|
253
|
+
await writer.drain()
|
|
254
|
+
|
|
255
|
+
# Read status line.
|
|
256
|
+
status_line = await reader.readline()
|
|
257
|
+
status_str = status_line.decode("utf-8", errors="replace").strip()
|
|
258
|
+
parts = status_str.split(" ", 2)
|
|
259
|
+
if len(parts) < 2:
|
|
260
|
+
raise NotFoundError(f"invalid HTTP response: {status_str}")
|
|
261
|
+
|
|
262
|
+
status_code = int(parts[1])
|
|
263
|
+
|
|
264
|
+
# Read headers (discard them, we just need the body stream).
|
|
265
|
+
while True:
|
|
266
|
+
header_line = await reader.readline()
|
|
267
|
+
if header_line in (b"\r\n", b"\n", b""):
|
|
268
|
+
break
|
|
269
|
+
|
|
270
|
+
if status_code == HTTPStatus.NOT_FOUND:
|
|
271
|
+
raise NotFoundError("model not found (HTTP 404)")
|
|
272
|
+
if status_code >= 400:
|
|
273
|
+
raise RunnerError(f"HTTP {status_code}")
|
|
274
|
+
|
|
275
|
+
return reader, writer
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
def _build_request_body(prompt: str, options: OllamaRunOptions) -> dict[str, Any]:
|
|
279
|
+
"""Build the JSON request body for POST /api/chat."""
|
|
280
|
+
messages: list[dict[str, str]] = []
|
|
281
|
+
|
|
282
|
+
system_prompt = options.system_prompt or ""
|
|
283
|
+
if options.append_system_prompt:
|
|
284
|
+
if system_prompt:
|
|
285
|
+
system_prompt += "\n" + options.append_system_prompt
|
|
286
|
+
else:
|
|
287
|
+
system_prompt = options.append_system_prompt
|
|
288
|
+
|
|
289
|
+
if system_prompt:
|
|
290
|
+
messages.append({"role": "system", "content": system_prompt})
|
|
291
|
+
|
|
292
|
+
messages.append({"role": "user", "content": prompt})
|
|
293
|
+
|
|
294
|
+
body: dict[str, Any] = {
|
|
295
|
+
"model": options.model,
|
|
296
|
+
"messages": messages,
|
|
297
|
+
"stream": True,
|
|
298
|
+
}
|
|
299
|
+
|
|
300
|
+
if options.think is not None:
|
|
301
|
+
body["think"] = options.think
|
|
302
|
+
if options.format:
|
|
303
|
+
body["format"] = options.format
|
|
304
|
+
if options.keep_alive:
|
|
305
|
+
body["keep_alive"] = options.keep_alive
|
|
306
|
+
|
|
307
|
+
model_opts = _build_model_options(options)
|
|
308
|
+
if model_opts:
|
|
309
|
+
body["options"] = model_opts
|
|
310
|
+
|
|
311
|
+
return body
|
|
312
|
+
|
|
313
|
+
|
|
314
|
+
def _build_model_options(options: OllamaRunOptions) -> dict[str, Any] | None:
|
|
315
|
+
"""Build the model options dict, or None if no options are set."""
|
|
316
|
+
opts = ModelOptions(
|
|
317
|
+
temperature=options.temperature,
|
|
318
|
+
num_ctx=options.num_ctx,
|
|
319
|
+
num_predict=options.num_predict,
|
|
320
|
+
seed=options.seed,
|
|
321
|
+
stop=options.stop,
|
|
322
|
+
top_k=options.top_k,
|
|
323
|
+
top_p=options.top_p,
|
|
324
|
+
min_p=options.min_p,
|
|
325
|
+
)
|
|
326
|
+
d = opts.to_dict()
|
|
327
|
+
return d if d else None
|