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.
@@ -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())