athena-claude-coder 0.4.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.
- athena_claude_coder/__init__.py +17 -0
- athena_claude_coder/app.py +492 -0
- athena_claude_coder/auth.py +45 -0
- athena_claude_coder/main.py +113 -0
- athena_claude_coder/permissions.py +222 -0
- athena_claude_coder/redact.py +75 -0
- athena_claude_coder/runner.py +1090 -0
- athena_claude_coder/schema.py +62 -0
- athena_claude_coder/settings.py +149 -0
- athena_claude_coder/store.py +650 -0
- athena_claude_coder/transcript.py +326 -0
- athena_claude_coder/worktrees.py +151 -0
- athena_claude_coder-0.4.1.dist-info/METADATA +308 -0
- athena_claude_coder-0.4.1.dist-info/RECORD +16 -0
- athena_claude_coder-0.4.1.dist-info/WHEEL +4 -0
- athena_claude_coder-0.4.1.dist-info/entry_points.txt +2 -0
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""In-guest Agent Protocol server for Athena coding tasks (ADR 0006).
|
|
2
|
+
|
|
3
|
+
The Athena deep agent delegates repository work to a ``claude-coder`` async
|
|
4
|
+
subagent. Its Agent Protocol calls arrive here — inside the tenant's Talos v2
|
|
5
|
+
computer, on loopback port 46100 — through the agora coding-worker proxy,
|
|
6
|
+
which is the authorization boundary. This package answers exactly the routes
|
|
7
|
+
the pinned ``langgraph_sdk`` client calls and drives one ``ClaudeSDKClient``
|
|
8
|
+
per thread against the checkout at ``/workspace/template``.
|
|
9
|
+
|
|
10
|
+
State lives on the persisted tree (``/workspace/.claude-state``): a SQLite
|
|
11
|
+
ledger of threads, runs, messages and Claude session ids, a per-run event
|
|
12
|
+
log, and the task's owner key as a single 0600 file that the ``apiKeyHelper``
|
|
13
|
+
script reads. Processes do not survive a computer suspend; the ledger does,
|
|
14
|
+
so a restart marks orphaned runs ``interrupted`` with the session id intact.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
__version__ = "0.4.1"
|
|
@@ -0,0 +1,492 @@
|
|
|
1
|
+
"""FastAPI Agent Protocol subset — exactly the routes langgraph_sdk 0.4.2 calls.
|
|
2
|
+
|
|
3
|
+
``deepagents.middleware.async_subagents`` drives ``threads.create``,
|
|
4
|
+
``runs.create``, ``runs.get``, ``threads.get`` (C6) and ``runs.cancel``;
|
|
5
|
+
``threads.get_state`` is kept for parity. ``PUT /internal/credential`` is the
|
|
6
|
+
proxy's key-rotation hook and ``GET /ok`` the unauthenticated health probe.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import asyncio
|
|
12
|
+
import contextlib
|
|
13
|
+
import logging
|
|
14
|
+
import uuid
|
|
15
|
+
from collections.abc import AsyncIterator
|
|
16
|
+
from contextlib import asynccontextmanager
|
|
17
|
+
from typing import Annotated, Any, Literal
|
|
18
|
+
|
|
19
|
+
from fastapi import APIRouter, Depends, FastAPI, Header, HTTPException, Query, Response
|
|
20
|
+
from pydantic import BaseModel, ConfigDict, Field
|
|
21
|
+
|
|
22
|
+
from athena_claude_coder.auth import (
|
|
23
|
+
GUEST_TOKEN_HEADER,
|
|
24
|
+
guest_token_dependency,
|
|
25
|
+
token_matches,
|
|
26
|
+
)
|
|
27
|
+
from athena_claude_coder.runner import (
|
|
28
|
+
HARD_MAX_CONCURRENT_TASKS,
|
|
29
|
+
ClientFactory,
|
|
30
|
+
PidAlive,
|
|
31
|
+
RunnerBusyError,
|
|
32
|
+
ThreadRunner,
|
|
33
|
+
clamp_slots,
|
|
34
|
+
)
|
|
35
|
+
from athena_claude_coder.settings import CoderSettings, load_settings
|
|
36
|
+
from athena_claude_coder.store import Store, ThreadExistsError, credential_name
|
|
37
|
+
from athena_claude_coder.transcript import (
|
|
38
|
+
DEFAULT_LIMIT,
|
|
39
|
+
find_transcript,
|
|
40
|
+
read_transcript,
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
logger = logging.getLogger(__name__)
|
|
44
|
+
|
|
45
|
+
SWEEP_EVERY_SECONDS = 600.0
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
ENVELOPE_MISSING_DETAIL = "run config lacks the athena credential envelope"
|
|
49
|
+
ENVELOPE_BASE_URL_DETAIL = "athena envelope lacks anthropic_base_url"
|
|
50
|
+
EMPTY_TASK_DETAIL = "empty task"
|
|
51
|
+
THREAD_NOT_FOUND_DETAIL = "thread not found"
|
|
52
|
+
THREAD_EXISTS_DETAIL = "thread already exists"
|
|
53
|
+
RUN_NOT_FOUND_DETAIL = "run not found"
|
|
54
|
+
NO_SESSION_DETAIL = "no session to hand off"
|
|
55
|
+
TRANSCRIPT_NOT_FOUND_DETAIL = "no transcript for this thread yet"
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
class ThreadCreateBody(BaseModel):
|
|
59
|
+
"""``POST /threads`` — what ``langgraph_sdk`` ``threads.create`` sends."""
|
|
60
|
+
|
|
61
|
+
model_config = ConfigDict(extra="ignore")
|
|
62
|
+
|
|
63
|
+
thread_id: str | None = None
|
|
64
|
+
metadata: dict[str, Any] | None = None
|
|
65
|
+
if_exists: Literal["raise", "do_nothing"] = "raise"
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
class HandoffRequestBody(BaseModel):
|
|
69
|
+
"""``POST /internal/handoff/request`` — from the ``athena-claude`` launcher."""
|
|
70
|
+
|
|
71
|
+
model_config = ConfigDict(extra="ignore")
|
|
72
|
+
|
|
73
|
+
thread_id: str | None = None
|
|
74
|
+
launcher_pid: int | None = Field(default=None, ge=1)
|
|
75
|
+
"""The launcher's PID; lets a lock whose terminal died be recognised as stale."""
|
|
76
|
+
takeover: bool = False
|
|
77
|
+
"""Re-mint the lease even though another launcher holds the session."""
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
class HandoffReturnBody(BaseModel):
|
|
81
|
+
"""``POST /internal/handoff/return`` — sent by the launcher when Claude exits."""
|
|
82
|
+
|
|
83
|
+
model_config = ConfigDict(extra="ignore")
|
|
84
|
+
|
|
85
|
+
thread_id: str | None = None
|
|
86
|
+
claude_session_id: str | None = None
|
|
87
|
+
owner_since: str | None = None
|
|
88
|
+
"""The lease the request handed out; a mismatch makes the return a no-op."""
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
class CredentialBody(BaseModel):
|
|
92
|
+
"""``PUT /internal/credential`` — the proxy's rotation push."""
|
|
93
|
+
|
|
94
|
+
model_config = ConfigDict(extra="ignore")
|
|
95
|
+
|
|
96
|
+
owner_key: str = Field(min_length=1)
|
|
97
|
+
expires_at: str | int | float | None = None
|
|
98
|
+
budget_user_id: str = ""
|
|
99
|
+
anthropic_base_url: str = Field(min_length=1)
|
|
100
|
+
thread_id: str | None = None
|
|
101
|
+
"""Also rotate this task's own credential file, not only the shared one."""
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def _task_text(run_input: Any) -> str:
|
|
105
|
+
"""Join the user-authored message content of a LangGraph ``input``.
|
|
106
|
+
|
|
107
|
+
Accepts ``{"role": "user", "content": str}`` (what deepagents sends),
|
|
108
|
+
LangChain-style ``{"type": "human", ...}``, and list-of-blocks content.
|
|
109
|
+
"""
|
|
110
|
+
if not isinstance(run_input, dict):
|
|
111
|
+
return ""
|
|
112
|
+
messages = run_input.get("messages")
|
|
113
|
+
if not isinstance(messages, list):
|
|
114
|
+
return ""
|
|
115
|
+
parts: list[str] = []
|
|
116
|
+
for message in messages:
|
|
117
|
+
if not isinstance(message, dict):
|
|
118
|
+
continue
|
|
119
|
+
kind = str(message.get("role") or message.get("type") or "user").lower()
|
|
120
|
+
if kind not in ("user", "human"):
|
|
121
|
+
continue
|
|
122
|
+
content = message.get("content", "")
|
|
123
|
+
if isinstance(content, str):
|
|
124
|
+
parts.append(content)
|
|
125
|
+
elif isinstance(content, list):
|
|
126
|
+
parts.extend(
|
|
127
|
+
str(block.get("text", ""))
|
|
128
|
+
for block in content
|
|
129
|
+
if isinstance(block, dict) and block.get("type", "text") == "text"
|
|
130
|
+
)
|
|
131
|
+
return "\n\n".join(part for part in parts if part)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
def _positive_int(raw: Any, *, default: int) -> int:
|
|
135
|
+
try:
|
|
136
|
+
value = int(raw)
|
|
137
|
+
except (TypeError, ValueError):
|
|
138
|
+
return default
|
|
139
|
+
return value if value >= 1 else default
|
|
140
|
+
|
|
141
|
+
|
|
142
|
+
def create_app(
|
|
143
|
+
settings: CoderSettings | None = None,
|
|
144
|
+
*,
|
|
145
|
+
client_factory: ClientFactory | None = None,
|
|
146
|
+
alive: PidAlive | None = None,
|
|
147
|
+
) -> FastAPI:
|
|
148
|
+
"""Build the server; `client_factory` and `alive` let tests substitute fakes."""
|
|
149
|
+
s = settings or load_settings()
|
|
150
|
+
store = Store(
|
|
151
|
+
sqlite_path=s.sqlite_path,
|
|
152
|
+
credential_path=s.credential_path,
|
|
153
|
+
credential_meta_path=s.credential_meta_path,
|
|
154
|
+
)
|
|
155
|
+
runner = ThreadRunner(
|
|
156
|
+
settings=s, store=store, client_factory=client_factory, alive=alive
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
@asynccontextmanager
|
|
160
|
+
async def lifespan(_: FastAPI) -> AsyncIterator[None]:
|
|
161
|
+
s.state_dir.mkdir(parents=True, exist_ok=True)
|
|
162
|
+
await store.open()
|
|
163
|
+
# Processes never survive a computer suspend; live-looking rows are orphans.
|
|
164
|
+
orphans = await store.mark_orphans_interrupted()
|
|
165
|
+
if orphans:
|
|
166
|
+
logger.warning("marked %d orphaned run(s) interrupted at startup", orphans)
|
|
167
|
+
stale = await runner.sweep_stale_human_locks()
|
|
168
|
+
if stale:
|
|
169
|
+
logger.warning("released stale human lock(s) at startup: %s", stale)
|
|
170
|
+
if not s.guest_token:
|
|
171
|
+
logger.warning(
|
|
172
|
+
"ATHENA_CODING_WORKER_TOKEN is not set; "
|
|
173
|
+
"every route except /ok answers 503"
|
|
174
|
+
)
|
|
175
|
+
sweep = asyncio.create_task(
|
|
176
|
+
runner.sweep_loop(every_seconds=SWEEP_EVERY_SECONDS),
|
|
177
|
+
name="claude-coder-sweep",
|
|
178
|
+
)
|
|
179
|
+
try:
|
|
180
|
+
yield
|
|
181
|
+
finally:
|
|
182
|
+
sweep.cancel()
|
|
183
|
+
with contextlib.suppress(asyncio.CancelledError):
|
|
184
|
+
await sweep
|
|
185
|
+
await runner.shutdown()
|
|
186
|
+
await store.close()
|
|
187
|
+
|
|
188
|
+
app = FastAPI(
|
|
189
|
+
title="athena-claude-coder",
|
|
190
|
+
version=s.package_version,
|
|
191
|
+
docs_url=None,
|
|
192
|
+
redoc_url=None,
|
|
193
|
+
openapi_url=None,
|
|
194
|
+
lifespan=lifespan,
|
|
195
|
+
)
|
|
196
|
+
app.state.settings = s
|
|
197
|
+
app.state.store = store
|
|
198
|
+
app.state.runner = runner
|
|
199
|
+
|
|
200
|
+
@app.get("/ok")
|
|
201
|
+
async def ok(
|
|
202
|
+
token: Annotated[str | None, Header(alias=GUEST_TOKEN_HEADER)] = None,
|
|
203
|
+
) -> dict[str, Any]:
|
|
204
|
+
# Unauthenticated on purpose — the ensure command probes it on loopback
|
|
205
|
+
# before the token exists — but never more than liveness for a caller
|
|
206
|
+
# without the token: whether a task is running, and on which thread,
|
|
207
|
+
# is for the proxy only. A wrong token still gets 200 here, not 401.
|
|
208
|
+
payload: dict[str, Any] = {"ok": True, "version": s.package_version}
|
|
209
|
+
if token_matches(expected=s.guest_token, presented=token):
|
|
210
|
+
slots = await runner.slot_state()
|
|
211
|
+
busy_thread = runner.busy_thread()
|
|
212
|
+
payload["busy"] = slots["busy"]
|
|
213
|
+
payload["active_thread_id"] = busy_thread
|
|
214
|
+
payload["active_thread_ids"] = slots["active_thread_ids"]
|
|
215
|
+
payload["slots"] = {"used": slots["used"], "total": slots["total"]}
|
|
216
|
+
# The owner lock of the thread that matters right now: the live one,
|
|
217
|
+
# else the human-held or most recently touched session — and which
|
|
218
|
+
# thread that is, so a human lock can be tied to its task once the
|
|
219
|
+
# handoff has drained the run and `active_thread_id` is gone.
|
|
220
|
+
session = (
|
|
221
|
+
await store.get_session(thread_id=busy_thread)
|
|
222
|
+
if busy_thread is not None
|
|
223
|
+
else await store.human_held_session() or await store.latest_session()
|
|
224
|
+
)
|
|
225
|
+
payload["owner"] = session["owner"] if session is not None else "none"
|
|
226
|
+
payload["owner_since"] = (
|
|
227
|
+
session["owner_since"] if session is not None else None
|
|
228
|
+
)
|
|
229
|
+
payload["owner_thread_id"] = (
|
|
230
|
+
session["thread_id"] if session is not None else None
|
|
231
|
+
)
|
|
232
|
+
return payload
|
|
233
|
+
|
|
234
|
+
protected = APIRouter(dependencies=[Depends(guest_token_dependency(s.guest_token))])
|
|
235
|
+
|
|
236
|
+
@protected.post("/threads")
|
|
237
|
+
async def create_thread(body: ThreadCreateBody | None = None) -> dict[str, Any]:
|
|
238
|
+
body = body or ThreadCreateBody()
|
|
239
|
+
thread_id = body.thread_id or str(uuid.uuid4())
|
|
240
|
+
try:
|
|
241
|
+
thread = await store.create_thread(
|
|
242
|
+
thread_id=thread_id,
|
|
243
|
+
metadata=body.metadata or {},
|
|
244
|
+
if_exists=body.if_exists,
|
|
245
|
+
)
|
|
246
|
+
except ThreadExistsError as exc:
|
|
247
|
+
raise HTTPException(status_code=409, detail=THREAD_EXISTS_DETAIL) from exc
|
|
248
|
+
logger.info("thread %s created", thread_id)
|
|
249
|
+
return dict(thread)
|
|
250
|
+
|
|
251
|
+
# (C6) deepagents reads the result through threads.get, NOT threads.get_state:
|
|
252
|
+
# GET /threads/{thread_id} returns the langgraph Thread shape with
|
|
253
|
+
# values.messages populated. Both routes exist; /state alone would 404 every
|
|
254
|
+
# check_async_task.
|
|
255
|
+
@protected.get("/threads/{thread_id}")
|
|
256
|
+
async def get_thread(thread_id: str) -> dict[str, Any]:
|
|
257
|
+
thread = await store.get_thread(thread_id=thread_id)
|
|
258
|
+
if thread is None:
|
|
259
|
+
raise HTTPException(status_code=404, detail=THREAD_NOT_FOUND_DETAIL)
|
|
260
|
+
return dict(thread)
|
|
261
|
+
|
|
262
|
+
@protected.get("/threads/{thread_id}/state")
|
|
263
|
+
async def get_state(thread_id: str) -> dict[str, Any]:
|
|
264
|
+
state = await store.get_thread_state(thread_id=thread_id)
|
|
265
|
+
if state is None:
|
|
266
|
+
raise HTTPException(status_code=404, detail=THREAD_NOT_FOUND_DETAIL)
|
|
267
|
+
return dict(state)
|
|
268
|
+
|
|
269
|
+
@protected.post("/threads/{thread_id}/runs")
|
|
270
|
+
async def create_run(thread_id: str, body: dict[str, Any]) -> dict[str, Any]:
|
|
271
|
+
if not await store.thread_exists(thread_id=thread_id):
|
|
272
|
+
raise HTTPException(status_code=404, detail=THREAD_NOT_FOUND_DETAIL)
|
|
273
|
+
# The envelope is injected by the proxy on every run create; it is the
|
|
274
|
+
# only channel for identity, credential and budget, never the prompt.
|
|
275
|
+
athena = ((body.get("config") or {}).get("configurable") or {}).get("athena")
|
|
276
|
+
if (
|
|
277
|
+
not isinstance(athena, dict)
|
|
278
|
+
or not str(athena.get("owner_key") or "").strip()
|
|
279
|
+
):
|
|
280
|
+
raise HTTPException(status_code=400, detail=ENVELOPE_MISSING_DETAIL)
|
|
281
|
+
anthropic_base_url = str(athena.get("anthropic_base_url") or "").strip()
|
|
282
|
+
if not anthropic_base_url:
|
|
283
|
+
raise HTTPException(status_code=400, detail=ENVELOPE_BASE_URL_DETAIL)
|
|
284
|
+
task_text = _task_text(body.get("input"))
|
|
285
|
+
if not task_text.strip():
|
|
286
|
+
raise HTTPException(status_code=400, detail=EMPTY_TASK_DETAIL)
|
|
287
|
+
strategy = str(body.get("multitask_strategy") or "reject")
|
|
288
|
+
# The workspace policy's slot count rides in the envelope; the guest owns
|
|
289
|
+
# the hard ceiling. More than one slot means a git worktree per thread.
|
|
290
|
+
requested_cap = _positive_int(athena.get("max_concurrent_tasks"), default=1)
|
|
291
|
+
max_concurrent_tasks = clamp_slots(requested_cap)
|
|
292
|
+
if requested_cap > max_concurrent_tasks:
|
|
293
|
+
logger.warning(
|
|
294
|
+
"envelope asks for %d concurrent tasks; "
|
|
295
|
+
"the hard maximum is %d — clamping",
|
|
296
|
+
requested_cap,
|
|
297
|
+
HARD_MAX_CONCURRENT_TASKS,
|
|
298
|
+
)
|
|
299
|
+
run_id = str(uuid.uuid4())
|
|
300
|
+
# Claim the slot before writing anything: a create that loses the race
|
|
301
|
+
# must leave no credential, transcript entry or run row behind.
|
|
302
|
+
try:
|
|
303
|
+
await runner.reserve(
|
|
304
|
+
thread_id=thread_id,
|
|
305
|
+
run_id=run_id,
|
|
306
|
+
multitask_strategy=strategy,
|
|
307
|
+
max_concurrent_tasks=max_concurrent_tasks,
|
|
308
|
+
)
|
|
309
|
+
except RunnerBusyError as exc:
|
|
310
|
+
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
|
311
|
+
try:
|
|
312
|
+
# The shared file (human launcher, baked helper) and this task's own
|
|
313
|
+
# file (what the run's settings point Claude at).
|
|
314
|
+
for name in (None, credential_name(thread_id)):
|
|
315
|
+
await store.put_credential(
|
|
316
|
+
owner_key=str(athena["owner_key"]),
|
|
317
|
+
expires_at=str(athena.get("owner_key_expires_at") or ""),
|
|
318
|
+
budget_user_id=str(athena.get("budget_user_id") or ""),
|
|
319
|
+
anthropic_base_url=anthropic_base_url,
|
|
320
|
+
name=name,
|
|
321
|
+
)
|
|
322
|
+
await store.append_message(
|
|
323
|
+
thread_id=thread_id, type="human", content=task_text
|
|
324
|
+
)
|
|
325
|
+
created = await store.create_run(
|
|
326
|
+
run_id=run_id, thread_id=thread_id, multitask_strategy=strategy
|
|
327
|
+
)
|
|
328
|
+
await runner.start(
|
|
329
|
+
thread_id=thread_id,
|
|
330
|
+
run_id=run_id,
|
|
331
|
+
task_text=task_text,
|
|
332
|
+
athena=athena,
|
|
333
|
+
multitask_strategy=strategy,
|
|
334
|
+
max_concurrent_tasks=max_concurrent_tasks,
|
|
335
|
+
)
|
|
336
|
+
except RunnerBusyError as exc:
|
|
337
|
+
# Unreachable while the reservation is honoured; kept so a regression
|
|
338
|
+
# surfaces as a 409 with an error row rather than a stray pending run.
|
|
339
|
+
runner.release(thread_id=thread_id, run_id=run_id)
|
|
340
|
+
await store.set_run_status(run_id=run_id, status="error", error=str(exc))
|
|
341
|
+
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
|
342
|
+
except BaseException:
|
|
343
|
+
runner.release(thread_id=thread_id, run_id=run_id)
|
|
344
|
+
raise
|
|
345
|
+
logger.info(
|
|
346
|
+
"run %s created on thread %s (strategy=%s, parent=%s)",
|
|
347
|
+
run_id,
|
|
348
|
+
thread_id,
|
|
349
|
+
strategy,
|
|
350
|
+
athena.get("thread_id"),
|
|
351
|
+
)
|
|
352
|
+
# The run task may already have flipped the row to `running`; report
|
|
353
|
+
# whatever the ledger says now.
|
|
354
|
+
run = await store.get_run(run_id=run_id) or created
|
|
355
|
+
return dict(run)
|
|
356
|
+
|
|
357
|
+
@protected.get("/threads/{thread_id}/runs/{run_id}")
|
|
358
|
+
async def get_run(thread_id: str, run_id: str) -> dict[str, Any]:
|
|
359
|
+
run = await store.get_run(run_id=run_id)
|
|
360
|
+
if run is None or run["thread_id"] != thread_id:
|
|
361
|
+
raise HTTPException(status_code=404, detail=RUN_NOT_FOUND_DETAIL)
|
|
362
|
+
return dict(run)
|
|
363
|
+
|
|
364
|
+
@protected.post("/threads/{thread_id}/runs/{run_id}/cancel", status_code=204)
|
|
365
|
+
async def cancel_run(
|
|
366
|
+
thread_id: str,
|
|
367
|
+
run_id: str,
|
|
368
|
+
wait: Annotated[bool, Query()] = False,
|
|
369
|
+
action: Annotated[str, Query()] = "interrupt",
|
|
370
|
+
) -> Response:
|
|
371
|
+
run = await store.get_run(run_id=run_id)
|
|
372
|
+
if run is None or run["thread_id"] != thread_id:
|
|
373
|
+
raise HTTPException(status_code=404, detail=RUN_NOT_FOUND_DETAIL)
|
|
374
|
+
# `action=rollback` would delete the run and its checkpoint on a
|
|
375
|
+
# LangGraph server; here every cancel is an interrupt — the workspace
|
|
376
|
+
# is the tenant's and is never deleted (proposal A2). `wait` is
|
|
377
|
+
# implicit: cancel returns after the bounded drain.
|
|
378
|
+
logger.info("run %s cancel (wait=%s, action=%s)", run_id, wait, action)
|
|
379
|
+
await runner.cancel(thread_id=thread_id, run_id=run_id)
|
|
380
|
+
return Response(status_code=204)
|
|
381
|
+
|
|
382
|
+
@protected.put("/internal/credential", status_code=204)
|
|
383
|
+
async def put_credential(body: CredentialBody) -> Response:
|
|
384
|
+
names: list[str | None] = [None]
|
|
385
|
+
if body.thread_id:
|
|
386
|
+
names.append(credential_name(body.thread_id))
|
|
387
|
+
for name in names:
|
|
388
|
+
await store.put_credential(
|
|
389
|
+
owner_key=body.owner_key,
|
|
390
|
+
expires_at="" if body.expires_at is None else str(body.expires_at),
|
|
391
|
+
budget_user_id=body.budget_user_id,
|
|
392
|
+
anthropic_base_url=body.anthropic_base_url,
|
|
393
|
+
name=name,
|
|
394
|
+
)
|
|
395
|
+
logger.info("credential rotated (budget_user=%s)", body.budget_user_id)
|
|
396
|
+
return Response(status_code=204)
|
|
397
|
+
|
|
398
|
+
async def _resolve_thread(requested: str | None, *, for_return: bool) -> str:
|
|
399
|
+
"""Explicit thread id, else the live thread, else the latest session."""
|
|
400
|
+
if requested:
|
|
401
|
+
if not await store.thread_exists(thread_id=requested):
|
|
402
|
+
raise HTTPException(status_code=404, detail=THREAD_NOT_FOUND_DETAIL)
|
|
403
|
+
return requested
|
|
404
|
+
busy_thread = runner.busy_thread()
|
|
405
|
+
if busy_thread is not None:
|
|
406
|
+
return busy_thread
|
|
407
|
+
session = (await store.human_held_session() if for_return else None) or (
|
|
408
|
+
await store.latest_session()
|
|
409
|
+
)
|
|
410
|
+
if session is None:
|
|
411
|
+
raise HTTPException(status_code=404, detail=NO_SESSION_DETAIL)
|
|
412
|
+
return session["thread_id"]
|
|
413
|
+
|
|
414
|
+
@protected.post("/internal/handoff/request")
|
|
415
|
+
async def handoff_request(body: HandoffRequestBody | None = None) -> dict[str, Any]:
|
|
416
|
+
body = body or HandoffRequestBody()
|
|
417
|
+
thread_id = await _resolve_thread(body.thread_id, for_return=False)
|
|
418
|
+
try:
|
|
419
|
+
result = await runner.handoff_request(
|
|
420
|
+
thread_id=thread_id,
|
|
421
|
+
launcher_pid=body.launcher_pid,
|
|
422
|
+
takeover=body.takeover,
|
|
423
|
+
)
|
|
424
|
+
except RunnerBusyError as exc:
|
|
425
|
+
raise HTTPException(status_code=409, detail=str(exc)) from exc
|
|
426
|
+
row = result["session"]
|
|
427
|
+
logger.info(
|
|
428
|
+
"thread %s: handoff request served (owner=%s, acquired=%s)",
|
|
429
|
+
thread_id,
|
|
430
|
+
row["owner"],
|
|
431
|
+
result["acquired"],
|
|
432
|
+
)
|
|
433
|
+
return {
|
|
434
|
+
"thread_id": row["thread_id"],
|
|
435
|
+
"claude_session_id": row["claude_session_id"],
|
|
436
|
+
"cwd": row["cwd"],
|
|
437
|
+
"owner": row["owner"],
|
|
438
|
+
"owner_since": row["owner_since"],
|
|
439
|
+
"acquired": result["acquired"],
|
|
440
|
+
}
|
|
441
|
+
|
|
442
|
+
@protected.post("/internal/handoff/return")
|
|
443
|
+
async def handoff_return(body: HandoffReturnBody | None = None) -> dict[str, Any]:
|
|
444
|
+
body = body or HandoffReturnBody()
|
|
445
|
+
thread_id = await _resolve_thread(body.thread_id, for_return=True)
|
|
446
|
+
row = await runner.handoff_return(
|
|
447
|
+
thread_id=thread_id,
|
|
448
|
+
claude_session_id=body.claude_session_id,
|
|
449
|
+
lease=body.owner_since,
|
|
450
|
+
)
|
|
451
|
+
logger.info(
|
|
452
|
+
"thread %s: handoff returned (session=%s)",
|
|
453
|
+
thread_id,
|
|
454
|
+
row["claude_session_id"],
|
|
455
|
+
)
|
|
456
|
+
return {"thread_id": row["thread_id"], "owner": row["owner"]}
|
|
457
|
+
|
|
458
|
+
# (P3.1) Agora reads a task's Claude turns through the guest — the transcript
|
|
459
|
+
# Claude Code writes under CLAUDE_CONFIG_DIR — because the session bridge has
|
|
460
|
+
# no credential a guest could hold. Read-only, never under the runner lock,
|
|
461
|
+
# parsed in a worker thread so a large file never blocks the server.
|
|
462
|
+
@protected.get("/threads/{thread_id}/transcript")
|
|
463
|
+
async def get_transcript(
|
|
464
|
+
thread_id: str,
|
|
465
|
+
after: Annotated[int, Query()] = 0,
|
|
466
|
+
limit: Annotated[int, Query()] = DEFAULT_LIMIT,
|
|
467
|
+
) -> dict[str, Any]:
|
|
468
|
+
if not await store.thread_exists(thread_id=thread_id):
|
|
469
|
+
raise HTTPException(status_code=404, detail=THREAD_NOT_FOUND_DETAIL)
|
|
470
|
+
session = await store.get_session(thread_id=thread_id)
|
|
471
|
+
claude_session_id = session["claude_session_id"] if session else None
|
|
472
|
+
if not claude_session_id:
|
|
473
|
+
raise HTTPException(status_code=404, detail=TRANSCRIPT_NOT_FOUND_DETAIL)
|
|
474
|
+
path = find_transcript(
|
|
475
|
+
config_dir=s.claude_config_dir,
|
|
476
|
+
session_id=claude_session_id,
|
|
477
|
+
cwd=session["cwd"] if session else None,
|
|
478
|
+
)
|
|
479
|
+
if path is None:
|
|
480
|
+
raise HTTPException(status_code=404, detail=TRANSCRIPT_NOT_FOUND_DETAIL)
|
|
481
|
+
page = await asyncio.to_thread(
|
|
482
|
+
read_transcript, path=path, after=after, limit=limit
|
|
483
|
+
)
|
|
484
|
+
return {
|
|
485
|
+
"thread_id": thread_id,
|
|
486
|
+
"claude_session_id": claude_session_id,
|
|
487
|
+
"total": page["total"],
|
|
488
|
+
"entries": page["entries"],
|
|
489
|
+
}
|
|
490
|
+
|
|
491
|
+
app.include_router(protected)
|
|
492
|
+
return app
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
"""The guest shared secret (C7).
|
|
2
|
+
|
|
3
|
+
Anyone who can mint a router grant for port 46100 reaches this server, so a
|
|
4
|
+
grant alone is not enough: the proxy's ensure command exports a random
|
|
5
|
+
``ATHENA_CODING_WORKER_TOKEN`` into the server's environment and the proxy
|
|
6
|
+
sends the same value on every request. ``GET /ok`` stays open — the ensure
|
|
7
|
+
command curls it on loopback before the token exists.
|
|
8
|
+
|
|
9
|
+
Fail closed: with no token configured, every protected route answers 503
|
|
10
|
+
rather than letting the comparison degrade to "anything matches empty".
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import hmac
|
|
16
|
+
from collections.abc import Awaitable, Callable
|
|
17
|
+
from typing import Annotated
|
|
18
|
+
|
|
19
|
+
from fastapi import Header, HTTPException
|
|
20
|
+
|
|
21
|
+
GUEST_TOKEN_HEADER = "X-Athena-Coding-Worker-Token"
|
|
22
|
+
|
|
23
|
+
TOKEN_NOT_CONFIGURED_DETAIL = "guest token not configured"
|
|
24
|
+
TOKEN_INVALID_DETAIL = "invalid guest token"
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def token_matches(*, expected: str | None, presented: str | None) -> bool:
|
|
28
|
+
"""Constant-time equality; ``False`` whenever either side is unset or empty."""
|
|
29
|
+
if not expected or not presented:
|
|
30
|
+
return False
|
|
31
|
+
return hmac.compare_digest(presented.encode("utf-8"), expected.encode("utf-8"))
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def guest_token_dependency(expected: str | None) -> Callable[..., Awaitable[None]]:
|
|
35
|
+
"""Build the FastAPI dependency that guards every route except ``/ok``."""
|
|
36
|
+
|
|
37
|
+
async def require_guest_token(
|
|
38
|
+
token: Annotated[str | None, Header(alias=GUEST_TOKEN_HEADER)] = None,
|
|
39
|
+
) -> None:
|
|
40
|
+
if not expected:
|
|
41
|
+
raise HTTPException(status_code=503, detail=TOKEN_NOT_CONFIGURED_DETAIL)
|
|
42
|
+
if not token_matches(expected=expected, presented=token):
|
|
43
|
+
raise HTTPException(status_code=401, detail=TOKEN_INVALID_DETAIL)
|
|
44
|
+
|
|
45
|
+
return require_guest_token
|
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
"""Console entry point: ``athena-claude-coder serve`` and ``--version``."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import argparse
|
|
6
|
+
import json
|
|
7
|
+
import logging
|
|
8
|
+
import os
|
|
9
|
+
from collections.abc import MutableMapping, Sequence
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
from athena_claude_coder import __version__
|
|
14
|
+
from athena_claude_coder.settings import CoderSettings, load_settings
|
|
15
|
+
from athena_claude_coder.store import write_private_file
|
|
16
|
+
|
|
17
|
+
logger = logging.getLogger(__name__)
|
|
18
|
+
|
|
19
|
+
SCRUBBED_ENV_VARS: tuple[str, ...] = (
|
|
20
|
+
"ANTHROPIC_API_KEY",
|
|
21
|
+
"ANTHROPIC_AUTH_TOKEN",
|
|
22
|
+
"CLAUDE_CODE_OAUTH_TOKEN",
|
|
23
|
+
"CLAUDE_CODE_API_KEY_FILE_DESCRIPTOR",
|
|
24
|
+
)
|
|
25
|
+
"""Credential variables the server must never hand to a Claude process (C9)."""
|
|
26
|
+
|
|
27
|
+
API_KEY_HELPER_PATH = "/opt/athena/bin/athena-claude-key"
|
|
28
|
+
"""Installed by the bake; prints ``<state_dir>/credential`` and nothing else."""
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def scrub_credential_env(environ: MutableMapping[str, str] | None = None) -> list[str]:
|
|
32
|
+
"""Remove every credential variable from `environ`; return the names removed.
|
|
33
|
+
|
|
34
|
+
(C9) ``ClaudeAgentOptions.env`` is merged *over* ``os.environ`` and cannot
|
|
35
|
+
unset a key. The guest boot env carries the per-computer gateway key and
|
|
36
|
+
pm2 hands it to this process, so without this pop every task would bill
|
|
37
|
+
the computer's uncapped key instead of its own task budget. The SDK reads
|
|
38
|
+
``os.environ`` at connect time, so scrubbing once here covers every client
|
|
39
|
+
the runner builds.
|
|
40
|
+
"""
|
|
41
|
+
env = os.environ if environ is None else environ
|
|
42
|
+
removed: list[str] = []
|
|
43
|
+
for name in SCRUBBED_ENV_VARS:
|
|
44
|
+
if env.pop(name, None) is not None:
|
|
45
|
+
removed.append(name)
|
|
46
|
+
return removed
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def claude_settings_document() -> dict[str, Any]:
|
|
50
|
+
"""The ``--settings`` file every Claude session is launched with."""
|
|
51
|
+
return {
|
|
52
|
+
"apiKeyHelper": API_KEY_HELPER_PATH,
|
|
53
|
+
"env": {"ENABLE_TOOL_SEARCH": "true"},
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def write_claude_settings(*, settings: CoderSettings) -> Path:
|
|
58
|
+
path = settings.claude_settings_path
|
|
59
|
+
write_private_file(
|
|
60
|
+
path=path, content=json.dumps(claude_settings_document(), indent=2)
|
|
61
|
+
)
|
|
62
|
+
return path
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def serve(*, settings: CoderSettings | None = None) -> None:
|
|
66
|
+
"""Scrub the inherited credential env, write settings.json, run uvicorn."""
|
|
67
|
+
scrubbed = scrub_credential_env()
|
|
68
|
+
logging.basicConfig(
|
|
69
|
+
level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s"
|
|
70
|
+
)
|
|
71
|
+
if scrubbed:
|
|
72
|
+
logger.info("scrubbed inherited credential env: %s", ", ".join(scrubbed))
|
|
73
|
+
s = settings or load_settings()
|
|
74
|
+
s.state_dir.mkdir(parents=True, exist_ok=True)
|
|
75
|
+
write_claude_settings(settings=s)
|
|
76
|
+
|
|
77
|
+
import uvicorn
|
|
78
|
+
|
|
79
|
+
from athena_claude_coder.app import create_app
|
|
80
|
+
|
|
81
|
+
logger.info(
|
|
82
|
+
"athena-claude-coder %s serving on %s:%d (state=%s, repo=%s)",
|
|
83
|
+
s.package_version,
|
|
84
|
+
s.bind_host,
|
|
85
|
+
s.port,
|
|
86
|
+
s.state_dir,
|
|
87
|
+
s.repo_dir,
|
|
88
|
+
)
|
|
89
|
+
uvicorn.run(create_app(s), host=s.bind_host, port=s.port, log_level="info")
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def build_parser() -> argparse.ArgumentParser:
|
|
93
|
+
parser = argparse.ArgumentParser(
|
|
94
|
+
prog="athena-claude-coder",
|
|
95
|
+
description="In-guest Agent Protocol server for Athena coding tasks",
|
|
96
|
+
)
|
|
97
|
+
parser.add_argument(
|
|
98
|
+
"--version", action="version", version=f"%(prog)s {__version__}"
|
|
99
|
+
)
|
|
100
|
+
subcommands = parser.add_subparsers(dest="command", required=True)
|
|
101
|
+
subcommands.add_parser("serve", help="run the server on the loopback port")
|
|
102
|
+
return parser
|
|
103
|
+
|
|
104
|
+
|
|
105
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
106
|
+
args = build_parser().parse_args(argv)
|
|
107
|
+
if args.command == "serve":
|
|
108
|
+
serve()
|
|
109
|
+
return 0
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
if __name__ == "__main__": # pragma: no cover
|
|
113
|
+
raise SystemExit(main())
|