pulse-coding-agent 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.
- pulse/__init__.py +5 -0
- pulse/__main__.py +4 -0
- pulse/agent.py +270 -0
- pulse/agent_manager.py +335 -0
- pulse/audit.py +70 -0
- pulse/auth.py +670 -0
- pulse/ci/github_client.py +66 -0
- pulse/ci/runner.py +28 -0
- pulse/cli.py +1075 -0
- pulse/cli_ui.py +977 -0
- pulse/config.py +167 -0
- pulse/context.py +960 -0
- pulse/conversations/__init__.py +8 -0
- pulse/conversations/manager.py +312 -0
- pulse/core/agent.py +188 -0
- pulse/core/planner.py +105 -0
- pulse/core/protocols.py +37 -0
- pulse/edits.py +65 -0
- pulse/episodic.py +93 -0
- pulse/eval/__init__.py +8 -0
- pulse/eval/trajectory_logger.py +91 -0
- pulse/eval/verifier.py +133 -0
- pulse/execution/__init__.py +5 -0
- pulse/execution/remote_task.py +76 -0
- pulse/git.py +162 -0
- pulse/interactive.py +234 -0
- pulse/mcp/__init__.py +4 -0
- pulse/mcp/client.py +215 -0
- pulse/mcp/local_tools.py +105 -0
- pulse/memory.py +212 -0
- pulse/mutations.py +283 -0
- pulse/orchestration/__init__.py +3 -0
- pulse/orchestration/orchestrator.py +162 -0
- pulse/patch.py +129 -0
- pulse/planner/__init__.py +3 -0
- pulse/planner/dag_planner.py +85 -0
- pulse/planner/execution_loop.py +159 -0
- pulse/production.py +235 -0
- pulse/provider.py +59 -0
- pulse/provider_keys.py +278 -0
- pulse/providers/__init__.py +26 -0
- pulse/providers/anthropic.py +65 -0
- pulse/providers/base.py +251 -0
- pulse/providers/deepseek.py +10 -0
- pulse/providers/failover.py +32 -0
- pulse/providers/gemini.py +66 -0
- pulse/providers/groq.py +10 -0
- pulse/providers/manager.py +262 -0
- pulse/providers/openai.py +40 -0
- pulse/providers/openrouter.py +20 -0
- pulse/py.typed +1 -0
- pulse/reasoning.py +570 -0
- pulse/refactor/__init__.py +3 -0
- pulse/refactor/impact_analyzer.py +44 -0
- pulse/repository.py +209 -0
- pulse/rpc.py +249 -0
- pulse/rule_synthesizer.py +54 -0
- pulse/runtime.py +217 -0
- pulse/safety/__init__.py +3 -0
- pulse/safety/safety_manager.py +97 -0
- pulse/sandbox/SECURITY.md +57 -0
- pulse/sandbox/__init__.py +57 -0
- pulse/sandbox/api.py +594 -0
- pulse/sandbox/audit.py +153 -0
- pulse/sandbox/backend/__init__.py +7 -0
- pulse/sandbox/backend/base.py +72 -0
- pulse/sandbox/backend/docker.py +498 -0
- pulse/sandbox/backend/host.py +140 -0
- pulse/sandbox/backend/remote.py +224 -0
- pulse/sandbox/errors.py +106 -0
- pulse/sandbox/filesystem.py +476 -0
- pulse/sandbox/git_safe.py +50 -0
- pulse/sandbox/lifecycle.py +88 -0
- pulse/sandbox/network.py +205 -0
- pulse/sandbox/path_validator.py +280 -0
- pulse/sandbox/policy.py +209 -0
- pulse/sandbox/process.py +331 -0
- pulse/sandbox/project.py +158 -0
- pulse/sandbox/python_safe.py +62 -0
- pulse/sandbox/remote/__init__.py +1 -0
- pulse/sandbox/remote/client.py +389 -0
- pulse/sandbox/remote/models.py +167 -0
- pulse/sandbox/remote/protocol.py +65 -0
- pulse/sandbox/remote/server.py +984 -0
- pulse/sandbox/remote/worker.py +175 -0
- pulse/sandbox/resources.py +236 -0
- pulse/sandbox/secrets.py +241 -0
- pulse/session_manager.py +365 -0
- pulse/software_engineer.py +189 -0
- pulse/storage.py +140 -0
- pulse/streaming.py +385 -0
- pulse/subprocesses.py +79 -0
- pulse/task_manager.py +2005 -0
- pulse/telemetry/__init__.py +25 -0
- pulse/telemetry/cost_tracker.py +95 -0
- pulse/telemetry/logger.py +110 -0
- pulse/tool_policy.py +197 -0
- pulse/tool_registry.py +163 -0
- pulse/tools.py +372 -0
- pulse/verification.py +118 -0
- pulse_coding_agent-0.1.0.dist-info/METADATA +211 -0
- pulse_coding_agent-0.1.0.dist-info/RECORD +104 -0
- pulse_coding_agent-0.1.0.dist-info/WHEEL +4 -0
- pulse_coding_agent-0.1.0.dist-info/entry_points.txt +4 -0
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
from pulse.core.protocols import StreamChunk
|
|
7
|
+
from pulse.providers.base import BaseProvider, ChatMessage
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class GeminiProvider(BaseProvider):
|
|
11
|
+
"""Google Gemini provider implementation for Pulse."""
|
|
12
|
+
|
|
13
|
+
api_key_env_var = "GEMINI_API_KEY"
|
|
14
|
+
|
|
15
|
+
@property
|
|
16
|
+
def endpoint(self) -> str:
|
|
17
|
+
model_name = self.config.name
|
|
18
|
+
if not model_name.startswith("models/"):
|
|
19
|
+
model_name = f"models/{model_name}"
|
|
20
|
+
return f"https://generativelanguage.googleapis.com/v1beta/{model_name}:streamGenerateContent"
|
|
21
|
+
|
|
22
|
+
def _headers(self) -> dict[str, str]:
|
|
23
|
+
return {
|
|
24
|
+
"x-goog-api-key": self.api_key or "",
|
|
25
|
+
"Content-Type": "application/json",
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
def _build_payload(
|
|
29
|
+
self,
|
|
30
|
+
messages: list[dict[str, Any] | ChatMessage],
|
|
31
|
+
temperature: float = 0.2,
|
|
32
|
+
) -> dict[str, Any]:
|
|
33
|
+
normalized = self._normalize_messages(messages)
|
|
34
|
+
contents: list[dict[str, Any]] = []
|
|
35
|
+
system_instruction = None
|
|
36
|
+
|
|
37
|
+
for msg in normalized:
|
|
38
|
+
role = msg.get("role", "user")
|
|
39
|
+
content = msg.get("content", "")
|
|
40
|
+
if role == "system":
|
|
41
|
+
system_instruction = {"parts": [{"text": str(content)}]}
|
|
42
|
+
else:
|
|
43
|
+
gemini_role = "model" if role == "assistant" else "user"
|
|
44
|
+
contents.append({"role": gemini_role, "parts": [{"text": str(content)}]})
|
|
45
|
+
|
|
46
|
+
payload: dict[str, Any] = {
|
|
47
|
+
"contents": contents,
|
|
48
|
+
"generationConfig": {
|
|
49
|
+
"temperature": temperature,
|
|
50
|
+
"maxOutputTokens": self.config.max_tokens,
|
|
51
|
+
},
|
|
52
|
+
}
|
|
53
|
+
if system_instruction:
|
|
54
|
+
payload["systemInstruction"] = system_instruction
|
|
55
|
+
return payload
|
|
56
|
+
|
|
57
|
+
def _parse_stream_chunk(self, payload_line: str) -> StreamChunk:
|
|
58
|
+
data = json.loads(payload_line)
|
|
59
|
+
candidates = data.get("candidates", [])
|
|
60
|
+
if not candidates:
|
|
61
|
+
return StreamChunk(content="", metadata={"raw": data})
|
|
62
|
+
|
|
63
|
+
parts = candidates[0].get("content", {}).get("parts", [])
|
|
64
|
+
text_parts = [part.get("text", "") for part in parts if isinstance(part, dict)]
|
|
65
|
+
content = "".join(text_parts)
|
|
66
|
+
return StreamChunk(content=content, metadata={"raw": data})
|
pulse/providers/groq.py
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from pulse.providers.openai import OpenAIProvider
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class GroqProvider(OpenAIProvider):
|
|
7
|
+
"""Groq provider implementation for Pulse."""
|
|
8
|
+
|
|
9
|
+
api_key_env_var = "GROQ_API_KEY"
|
|
10
|
+
endpoint = "https://api.groq.com/openai/v1/chat/completions"
|
|
@@ -0,0 +1,262 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
import tempfile
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from typing import Any
|
|
9
|
+
|
|
10
|
+
from pulse.config import ModelConfig
|
|
11
|
+
from pulse.providers.anthropic import AnthropicProvider
|
|
12
|
+
from pulse.providers.base import BaseProvider
|
|
13
|
+
from pulse.providers.deepseek import DeepSeekProvider
|
|
14
|
+
from pulse.providers.gemini import GeminiProvider
|
|
15
|
+
from pulse.providers.groq import GroqProvider
|
|
16
|
+
from pulse.providers.openai import OpenAIProvider
|
|
17
|
+
from pulse.providers.openrouter import OpenRouterProvider
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True)
|
|
21
|
+
class ModelMetadata:
|
|
22
|
+
name: str
|
|
23
|
+
speed: str # "Fast", "Balanced", "High Quality", "Ultra-Fast"
|
|
24
|
+
context_length: str # "128k", "200k", "1M", "2M", etc.
|
|
25
|
+
best_for: str # "Coding", "Reasoning", "General", "Vision"
|
|
26
|
+
category: str # "Flagship", "Coding", "Reasoning", "Fast"
|
|
27
|
+
status: str = "Active" # "Active", "Deprecated", "Beta"
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass(frozen=True)
|
|
31
|
+
class ProviderSpec:
|
|
32
|
+
key: str
|
|
33
|
+
display_name: str
|
|
34
|
+
env_var: str
|
|
35
|
+
provider_class: type[BaseProvider]
|
|
36
|
+
default_model: str
|
|
37
|
+
available_models: list[ModelMetadata]
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
PROVIDER_SPECS: dict[str, ProviderSpec] = {
|
|
41
|
+
"gemini": ProviderSpec(
|
|
42
|
+
key="gemini",
|
|
43
|
+
display_name="Google Gemini",
|
|
44
|
+
env_var="GEMINI_API_KEY",
|
|
45
|
+
provider_class=GeminiProvider,
|
|
46
|
+
default_model="gemini-2.0-flash",
|
|
47
|
+
available_models=[
|
|
48
|
+
ModelMetadata("gemini-2.0-flash", "Fast", "1M", "General & Fast Coding", "Flagship"),
|
|
49
|
+
ModelMetadata("gemini-1.5-pro", "High Quality", "2M", "Reasoning & Deep Analysis", "Reasoning"),
|
|
50
|
+
ModelMetadata("gemini-1.5-flash", "Ultra-Fast", "1M", "Lightweight Coding & Speed", "Fast"),
|
|
51
|
+
],
|
|
52
|
+
),
|
|
53
|
+
"openrouter": ProviderSpec(
|
|
54
|
+
key="openrouter",
|
|
55
|
+
display_name="OpenRouter",
|
|
56
|
+
env_var="OPENROUTER_API_KEY",
|
|
57
|
+
provider_class=OpenRouterProvider,
|
|
58
|
+
default_model="qwen/qwen3-coder",
|
|
59
|
+
available_models=[
|
|
60
|
+
ModelMetadata("qwen/qwen3-coder", "Balanced", "128k", "Advanced Coding & Refactoring", "Coding"),
|
|
61
|
+
ModelMetadata("anthropic/claude-3.5-sonnet", "High Quality", "200k", "Architecture & Technical Writing", "Flagship"),
|
|
62
|
+
ModelMetadata("deepseek/deepseek-r1", "High Quality", "164k", "Complex Reasoning & STEM", "Reasoning"),
|
|
63
|
+
ModelMetadata("google/gemini-2.0-flash-001", "Fast", "1M", "Fast Code Generation", "Fast"),
|
|
64
|
+
ModelMetadata("meta-llama/llama-3.3-70b-instruct", "Balanced", "128k", "General & Open Source", "General"),
|
|
65
|
+
],
|
|
66
|
+
),
|
|
67
|
+
"openai": ProviderSpec(
|
|
68
|
+
key="openai",
|
|
69
|
+
display_name="OpenAI",
|
|
70
|
+
env_var="OPENAI_API_KEY",
|
|
71
|
+
provider_class=OpenAIProvider,
|
|
72
|
+
default_model="gpt-4o",
|
|
73
|
+
available_models=[
|
|
74
|
+
ModelMetadata("gpt-4o", "High Quality", "128k", "Multimodal, Architecture & Coding", "Flagship"),
|
|
75
|
+
ModelMetadata("gpt-4o-mini", "Fast", "128k", "Lightweight Code & Fast Chat", "Fast"),
|
|
76
|
+
ModelMetadata("o1", "High Quality", "200k", "STEM & Complex Reasoning", "Reasoning"),
|
|
77
|
+
ModelMetadata("o3-mini", "Fast", "200k", "Fast Technical Reasoning & Math", "Reasoning"),
|
|
78
|
+
],
|
|
79
|
+
),
|
|
80
|
+
"anthropic": ProviderSpec(
|
|
81
|
+
key="anthropic",
|
|
82
|
+
display_name="Anthropic",
|
|
83
|
+
env_var="ANTHROPIC_API_KEY",
|
|
84
|
+
provider_class=AnthropicProvider,
|
|
85
|
+
default_model="claude-3-5-sonnet-20241022",
|
|
86
|
+
available_models=[
|
|
87
|
+
ModelMetadata("claude-3-5-sonnet-20241022", "High Quality", "200k", "State-of-the-Art Coding & Design", "Flagship"),
|
|
88
|
+
ModelMetadata("claude-3-5-haiku-20241022", "Fast", "200k", "Rapid Refactoring & Lightweight Tasks", "Fast"),
|
|
89
|
+
ModelMetadata("claude-3-opus-20240229", "High Quality", "200k", "Complex Analysis & System Architecture", "Reasoning"),
|
|
90
|
+
],
|
|
91
|
+
),
|
|
92
|
+
"groq": ProviderSpec(
|
|
93
|
+
key="groq",
|
|
94
|
+
display_name="Groq",
|
|
95
|
+
env_var="GROQ_API_KEY",
|
|
96
|
+
provider_class=GroqProvider,
|
|
97
|
+
default_model="llama-3.3-70b-versatile",
|
|
98
|
+
available_models=[
|
|
99
|
+
ModelMetadata("llama-3.3-70b-versatile", "Ultra-Fast", "128k", "General & Fast Coding", "Flagship"),
|
|
100
|
+
ModelMetadata("llama-3.1-8b-instant", "Ultra-Fast", "128k", "Instant Search & Micro-Edits", "Fast"),
|
|
101
|
+
ModelMetadata("mixtral-8x7b-32768", "Fast", "32k", "Fast Instruction Following", "General"),
|
|
102
|
+
ModelMetadata("deepseek-r1-distill-llama-70b", "Fast", "128k", "Fast STEM Reasoning", "Reasoning"),
|
|
103
|
+
],
|
|
104
|
+
),
|
|
105
|
+
"deepseek": ProviderSpec(
|
|
106
|
+
key="deepseek",
|
|
107
|
+
display_name="DeepSeek",
|
|
108
|
+
env_var="DEEPSEEK_API_KEY",
|
|
109
|
+
provider_class=DeepSeekProvider,
|
|
110
|
+
default_model="deepseek-chat",
|
|
111
|
+
available_models=[
|
|
112
|
+
ModelMetadata("deepseek-chat", "Balanced", "64k", "General Assistant & Coding (V3)", "Flagship"),
|
|
113
|
+
ModelMetadata("deepseek-reasoner", "High Quality", "64k", "Chain-of-Thought Reasoning (R1)", "Reasoning"),
|
|
114
|
+
],
|
|
115
|
+
),
|
|
116
|
+
}
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
class ProviderManager:
|
|
120
|
+
"""Central manager for Pulse single-active-model provider selection & configuration."""
|
|
121
|
+
|
|
122
|
+
def __init__(self, workspace: Path) -> None:
|
|
123
|
+
self.workspace = workspace.resolve()
|
|
124
|
+
self.config_file = self.workspace / ".agent" / "provider.json"
|
|
125
|
+
try:
|
|
126
|
+
self.config_file.resolve().relative_to(self.workspace)
|
|
127
|
+
except ValueError as error:
|
|
128
|
+
raise ValueError("Provider configuration must remain inside the workspace.") from error
|
|
129
|
+
|
|
130
|
+
def list_providers(self) -> list[dict[str, Any]]:
|
|
131
|
+
from pulse.provider_keys import ProviderKeyStore
|
|
132
|
+
|
|
133
|
+
statuses = {
|
|
134
|
+
status.provider: status
|
|
135
|
+
for status in ProviderKeyStore(self.workspace).statuses()
|
|
136
|
+
}
|
|
137
|
+
result = []
|
|
138
|
+
for spec in PROVIDER_SPECS.values():
|
|
139
|
+
result.append(
|
|
140
|
+
{
|
|
141
|
+
"key": spec.key,
|
|
142
|
+
"display_name": spec.display_name,
|
|
143
|
+
"env_var": spec.env_var,
|
|
144
|
+
"configured": statuses[spec.key].configured,
|
|
145
|
+
"key_source": statuses[spec.key].source,
|
|
146
|
+
"default_model": spec.default_model,
|
|
147
|
+
"models": spec.available_models,
|
|
148
|
+
}
|
|
149
|
+
)
|
|
150
|
+
return result
|
|
151
|
+
|
|
152
|
+
def get_provider_spec(self, provider_name: str) -> ProviderSpec:
|
|
153
|
+
normalized = provider_name.lower().strip()
|
|
154
|
+
if normalized not in PROVIDER_SPECS:
|
|
155
|
+
supported = ", ".join(PROVIDER_SPECS.keys())
|
|
156
|
+
raise ValueError(
|
|
157
|
+
f"Unsupported provider: '{provider_name}'. Supported providers: {supported}"
|
|
158
|
+
)
|
|
159
|
+
return PROVIDER_SPECS[normalized]
|
|
160
|
+
|
|
161
|
+
def get_model_metadata(
|
|
162
|
+
self, provider_name: str, model_name: str
|
|
163
|
+
) -> ModelMetadata | None:
|
|
164
|
+
try:
|
|
165
|
+
spec = self.get_provider_spec(provider_name)
|
|
166
|
+
for m in spec.available_models:
|
|
167
|
+
if m.name.lower() == model_name.lower():
|
|
168
|
+
return m
|
|
169
|
+
return ModelMetadata(
|
|
170
|
+
name=model_name,
|
|
171
|
+
speed="Custom",
|
|
172
|
+
context_length="128k",
|
|
173
|
+
best_for="User-Specified Model",
|
|
174
|
+
category="Custom",
|
|
175
|
+
)
|
|
176
|
+
except ValueError:
|
|
177
|
+
return None
|
|
178
|
+
|
|
179
|
+
def get_active_selection(self) -> tuple[str, str]:
|
|
180
|
+
"""Return (provider_name, model_name) stored in .agent/provider.json or default."""
|
|
181
|
+
if self.config_file.exists():
|
|
182
|
+
if self.config_file.is_symlink() or self.config_file.stat().st_size > 1_048_576:
|
|
183
|
+
return "openrouter", PROVIDER_SPECS["openrouter"].default_model
|
|
184
|
+
try:
|
|
185
|
+
data = json.loads(self.config_file.read_text(encoding="utf-8"))
|
|
186
|
+
provider = data.get("provider")
|
|
187
|
+
model = data.get("model")
|
|
188
|
+
if provider and model and provider.lower() in PROVIDER_SPECS:
|
|
189
|
+
return provider.lower(), model
|
|
190
|
+
# Intentionally broad to isolate execution boundaries and prevent crashes.
|
|
191
|
+
except Exception: # noqa: BLE001, S110
|
|
192
|
+
pass
|
|
193
|
+
return "openrouter", PROVIDER_SPECS["openrouter"].default_model
|
|
194
|
+
|
|
195
|
+
def validate_active_selection(self) -> tuple[str, str, str | None]:
|
|
196
|
+
"""Check active selection and fallback gracefully if model is deprecated or unsupported."""
|
|
197
|
+
provider, model = self.get_active_selection()
|
|
198
|
+
try:
|
|
199
|
+
spec = self.get_provider_spec(provider)
|
|
200
|
+
except ValueError:
|
|
201
|
+
default_spec = PROVIDER_SPECS["openrouter"]
|
|
202
|
+
self.save_selection("openrouter", default_spec.default_model)
|
|
203
|
+
return "openrouter", default_spec.default_model, f"Unknown provider '{provider}'. Switched to openrouter:{default_spec.default_model}."
|
|
204
|
+
|
|
205
|
+
meta = self.get_model_metadata(provider, model)
|
|
206
|
+
if meta and meta.status.lower() == "deprecated":
|
|
207
|
+
warning = f"Selected model '{model}' is deprecated for provider '{provider}'. Falling back to default model '{spec.default_model}'."
|
|
208
|
+
self.save_selection(provider, spec.default_model)
|
|
209
|
+
return provider, spec.default_model, warning
|
|
210
|
+
|
|
211
|
+
return provider, model, None
|
|
212
|
+
|
|
213
|
+
def save_selection(
|
|
214
|
+
self, provider_name: str, model_name: str | None = None
|
|
215
|
+
) -> tuple[str, str]:
|
|
216
|
+
spec = self.get_provider_spec(provider_name)
|
|
217
|
+
selected_model = (model_name or spec.default_model).strip()
|
|
218
|
+
if not selected_model or len(selected_model) > 256 or any(
|
|
219
|
+
ord(character) < 32 for character in selected_model
|
|
220
|
+
):
|
|
221
|
+
raise ValueError("Model identifiers must be 1-256 characters without controls.")
|
|
222
|
+
|
|
223
|
+
self.config_file.parent.mkdir(parents=True, exist_ok=True)
|
|
224
|
+
if self.config_file.parent.is_symlink() or self.config_file.is_symlink():
|
|
225
|
+
raise ValueError("Refusing to write provider configuration through a symbolic link.")
|
|
226
|
+
data = {
|
|
227
|
+
"schema_version": 1,
|
|
228
|
+
"provider": spec.key,
|
|
229
|
+
"model": selected_model,
|
|
230
|
+
}
|
|
231
|
+
temporary_name: str | None = None
|
|
232
|
+
try:
|
|
233
|
+
with tempfile.NamedTemporaryFile(
|
|
234
|
+
mode="w",
|
|
235
|
+
encoding="utf-8",
|
|
236
|
+
newline="\n",
|
|
237
|
+
prefix=".provider-",
|
|
238
|
+
suffix=".tmp",
|
|
239
|
+
dir=self.config_file.parent,
|
|
240
|
+
delete=False,
|
|
241
|
+
) as temporary:
|
|
242
|
+
temporary_name = temporary.name
|
|
243
|
+
json.dump(data, temporary, indent=2)
|
|
244
|
+
temporary.write("\n")
|
|
245
|
+
os.replace(temporary_name, self.config_file)
|
|
246
|
+
temporary_name = None
|
|
247
|
+
finally:
|
|
248
|
+
if temporary_name:
|
|
249
|
+
Path(temporary_name).unlink(missing_ok=True)
|
|
250
|
+
return spec.key, selected_model
|
|
251
|
+
|
|
252
|
+
def create_provider(
|
|
253
|
+
self,
|
|
254
|
+
config: ModelConfig,
|
|
255
|
+
workspace_env_path: Path | str | None = None,
|
|
256
|
+
api_key: str | None = None,
|
|
257
|
+
) -> BaseProvider:
|
|
258
|
+
env_path = workspace_env_path or (self.workspace / ".env")
|
|
259
|
+
spec = self.get_provider_spec(config.provider)
|
|
260
|
+
return spec.provider_class(
|
|
261
|
+
config=config, workspace_env_path=env_path, api_key=api_key
|
|
262
|
+
)
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
from pulse.core.protocols import StreamChunk
|
|
7
|
+
from pulse.providers.base import BaseProvider
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class OpenAIProvider(BaseProvider):
|
|
11
|
+
"""OpenAI provider implementation for Pulse."""
|
|
12
|
+
|
|
13
|
+
api_key_env_var = "OPENAI_API_KEY"
|
|
14
|
+
endpoint = "https://api.openai.com/v1/chat/completions"
|
|
15
|
+
|
|
16
|
+
def _build_payload(
|
|
17
|
+
self,
|
|
18
|
+
messages: list[dict[str, Any]],
|
|
19
|
+
temperature: float = 0.2,
|
|
20
|
+
) -> dict[str, Any]:
|
|
21
|
+
payload = super()._build_payload(messages, temperature)
|
|
22
|
+
if self.config.provider == "openai":
|
|
23
|
+
payload["stream_options"] = {"include_usage": True}
|
|
24
|
+
return payload
|
|
25
|
+
|
|
26
|
+
def _headers(self) -> dict[str, str]:
|
|
27
|
+
return {
|
|
28
|
+
"Authorization": f"Bearer {self.api_key or ''}",
|
|
29
|
+
"Content-Type": "application/json",
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
def _parse_stream_chunk(self, payload_line: str) -> StreamChunk:
|
|
33
|
+
data = json.loads(payload_line)
|
|
34
|
+
choices = data.get("choices", [])
|
|
35
|
+
if not choices:
|
|
36
|
+
return StreamChunk(content="", metadata={"raw": data})
|
|
37
|
+
|
|
38
|
+
delta = choices[0].get("delta", {})
|
|
39
|
+
content = delta.get("content", "")
|
|
40
|
+
return StreamChunk(content=content or "", metadata={"raw": data})
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from pulse.providers.openai import OpenAIProvider
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class OpenRouterProvider(OpenAIProvider):
|
|
7
|
+
"""OpenRouter provider implementation for Pulse."""
|
|
8
|
+
|
|
9
|
+
api_key_env_var = "OPENROUTER_API_KEY"
|
|
10
|
+
endpoint = "https://openrouter.ai/api/v1/chat/completions"
|
|
11
|
+
|
|
12
|
+
def _headers(self) -> dict[str, str]:
|
|
13
|
+
headers = super()._headers()
|
|
14
|
+
headers.update(
|
|
15
|
+
{
|
|
16
|
+
"HTTP-Referer": "https://local.pulse",
|
|
17
|
+
"X-Title": "Pulse",
|
|
18
|
+
}
|
|
19
|
+
)
|
|
20
|
+
return headers
|
pulse/py.typed
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
|