Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 7 additions & 2 deletions src/teamwork/routers/external.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from typing import Any

from fastapi import APIRouter, Depends, HTTPException, Header, Request
from pydantic import BaseModel
from pydantic import BaseModel, Field
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession

Expand Down Expand Up @@ -1096,6 +1096,10 @@ class ApprovalAsk(BaseModel):
project_id: str | None = None
payload: dict[str, Any] | None = None
reason: str | None = None
# How long the request stays answerable. Default is the model's one hour;
# an agent parking an unattended task (asked at 3am, answered at 8) asks
# for longer. Capped at a week.
expires_in_seconds: int | None = Field(default=None, ge=60, le=7 * 24 * 3600)


class ApprovalSpend(BaseModel):
Expand Down Expand Up @@ -1141,7 +1145,8 @@ async def ask_approval(
req = await request_for_client(
db, capability=body.capability, client_name=api_key.name,
agent_id=api_key.agent_id, project_id=body.project_id,
payload=body.payload, reason=body.reason)
payload=body.payload, reason=body.reason,
expires_in_seconds=body.expires_in_seconds)
await db.commit()
return _status_body(req)

Expand Down
11 changes: 10 additions & 1 deletion src/teamwork/services/approvals.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ async def request_approval(
db: AsyncSession, *, capability: str, client_name: str,
agent_id: str | None = None, project_id: str | None = None,
payload: dict[str, Any] | None = None, reason: str | None = None,
expires_in_seconds: int | None = None,
) -> ApprovalRequest:
"""Record a proposed action awaiting a decision.

Expand All @@ -47,14 +48,20 @@ async def request_approval(
ApprovalRequest.requested_by_client == client_name,
).order_by(ApprovalRequest.created_at.desc()).limit(1)
)).scalar_one_or_none()
wanted = (_now() + timedelta(seconds=expires_in_seconds)) if expires_in_seconds else None
if existing is not None and not existing.is_expired():
if wanted is not None and wanted > existing.expires_at:
existing.expires_at = wanted
await db.flush()
return existing

req = ApprovalRequest(
project_id=project_id, requested_by_agent_id=agent_id,
requested_by_client=client_name, capability=capability,
action_fingerprint=fingerprint, payload_preview=payload, reason=reason,
)
if wanted is not None:
req.expires_at = wanted
db.add(req)
await db.flush()
await append_event(
Expand Down Expand Up @@ -162,6 +169,7 @@ async def request_for_client(
db: AsyncSession, *, capability: str, client_name: str,
agent_id: str | None = None, project_id: str | None = None,
payload: dict[str, Any] | None = None, reason: str | None = None,
expires_in_seconds: int | None = None,
) -> ApprovalRequest:
"""An agent asking for approval of its own next action.

Expand All @@ -171,7 +179,8 @@ async def request_for_client(
"""
req = await request_approval(
db, capability=capability, client_name=client_name, agent_id=agent_id,
project_id=project_id, payload=payload, reason=reason)
project_id=project_id, payload=payload, reason=reason,
expires_in_seconds=expires_in_seconds)
if req.status != STATUS_PENDING:
return req
grant = await active_grant(db, client_name=client_name, capability=capability,
Expand Down
27 changes: 27 additions & 0 deletions tests/test_agent_self_approvals.py
Original file line number Diff line number Diff line change
Expand Up @@ -158,3 +158,30 @@ def test_with_a_ui_key_only_a_logged_in_session_decides(client, monkeypatch):
token, _ = issue_session_token("ui-key")
client.cookies.set(COOKIE, token)
assert _human_decide(client, aid).status_code == 200


def _hours_left(body) -> float:
from datetime import datetime
exp = datetime.fromisoformat(body["expires_at"])
return (exp - datetime.utcnow()).total_seconds() / 3600


def test_a_request_can_ask_to_stay_answerable_longer(client, monkeypatch):
# An unattended task that asks at 3am must still be answerable at 8am.
body = _ask(client, monkeypatch, expires_in_seconds=12 * 3600)
assert 11.9 < _hours_left(body) <= 12.0
# Asking again for the same action keeps one request and never shortens it.
again = _ask(client, monkeypatch, expires_in_seconds=3600)
assert again["approval_id"] == body["approval_id"]
assert _hours_left(again) > 11.9


def test_default_lifetime_is_unchanged(client, monkeypatch):
assert 0.9 < _hours_left(_ask(client, monkeypatch)) <= 1.0


def test_lifetime_is_bounded(client, monkeypatch):
_as(monkeypatch)
for bad in (5, 8 * 24 * 3600):
resp = client.post("/api/external/approvals", json={**ASK, "expires_in_seconds": bad})
assert resp.status_code == 422
Loading