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
111 changes: 85 additions & 26 deletions backend/druks/harnesses/providers.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,8 @@
# long — enough to authorize and paste, short enough that an abandoned attempt
# clears.
_CONNECT_PENDING_TTL_SECONDS = 600
# OpenAI expires a device code after 15 minutes.
_DEVICE_CODE_TTL_SECONDS = 15 * 60
# How often a fetch asks again while another refresher holds the row lock.
_LOCK_POLL_SECONDS = 0.5

Expand Down Expand Up @@ -92,9 +94,9 @@ def get_secret(cls, key: str) -> Secret:
raise NotImplementedError

@classmethod
async def connect_start(cls, *, account_id: str | None = None) -> tuple[str, str]:
"""Mint PKCE state under a single-use flow id; return (authorize URL,
flow id). A flow started by a resolved operator binds ``account_id``."""
async def start_connection(cls, *, account_id: str | None = None) -> dict:
"""Mint PKCE state under a single-use flow id; return the challenge the
operator answers. A flow started by a resolved operator binds ``account_id``."""
verifier = _b64url(secrets.token_bytes(64))
challenge = _b64url(hashlib.sha256(verifier.encode()).digest())
url, state = cls.authorize_url(verifier=verifier, challenge=challenge)
Expand All @@ -109,10 +111,10 @@ async def connect_start(cls, *, account_id: str | None = None) -> tuple[str, str
await get_client().set(
f"{CONNECT_PENDING_PREFIX}{flow_id}", pending, ex=_CONNECT_PENDING_TTL_SECONDS
)
return url, flow_id
return {"method": "code", "connection_id": flow_id, "authorize_url": url}

@classmethod
async def connect_complete(cls, *, flow_id: str, pasted: str) -> CompletedConnect:
async def complete_connection(cls, *, flow_id: str, pasted: str) -> CompletedConnect:
"""Pop the flow's single-use state, parse the paste, exchange the code.
Raises :class:`ConnectError` on failure; the state is gone either way,
so a retry re-starts cleanly."""
Expand All @@ -129,7 +131,20 @@ async def connect_complete(cls, *, flow_id: str, pasted: str) -> CompletedConnec
"That code is from a different connect attempt — start it again."
)

payload, provider_email = await cls.exchange(code=code, verifier=expected["verifier"])
return await cls._complete(
code=code, verifier=expected["verifier"], account_id=expected["account_id"]
)

@classmethod
async def check_connection(cls, *, flow_id: str) -> CompletedConnect | None:
"""The connect of a device flow after the operator approves it. None until then."""
raise exceptions.ConnectError(f"{cls.label} does not connect with a device code.")

@classmethod
async def _complete(
cls, *, code: str, verifier: str, account_id: str | None
) -> CompletedConnect:
payload, provider_email = await cls.exchange(code=code, verifier=verifier)
if not provider_email:
raise exceptions.ConnectError(
"The provider returned no account email — authorize with an account "
Expand All @@ -140,13 +155,13 @@ async def connect_complete(cls, *, flow_id: str, pasted: str) -> CompletedConnec
payload=payload,
provider_email=provider_email,
expires_at=expires_at,
account_id=expected["account_id"],
account_id=account_id,
)

@classmethod
def authorize_url(cls, *, verifier: str, challenge: str) -> tuple[str, str]:
"""Build this provider's PKCE authorize URL; return (url, state), where
``state`` is what the provider echoes back so connect_complete can
``state`` is what the provider echoes back so complete_connection can
verify the round-trip."""
raise NotImplementedError

Expand Down Expand Up @@ -735,7 +750,7 @@ def _grant_body(cls, refresh_token: str) -> dict:
@classmethod
def authorize_url(cls, *, verifier: str, challenge: str) -> tuple[str, str]:
# Anthropic's console flow echoes the PKCE verifier back as the OAuth
# state, so that's what connect_complete checks.
# state, so that's what complete_connection checks.
params = {
"code": "true",
"client_id": cls._CLIENT_ID,
Expand Down Expand Up @@ -923,9 +938,9 @@ class OpenAiProvider(Provider):
REFRESH_MARGIN = timedelta(hours=24)
_TOKEN_URL = "https://auth.openai.com/oauth/token"
_CLIENT_ID = "app_EMoamEEZ73f0CkXaXp7hrann"
# Connect-flow (PKCE): authorize on auth.openai.com; the operator pastes the
# failed localhost redirect URL back.
redirect_uri = "http://localhost:1455/auth/callback"
# Connect-flow (device code): the operator enters a code at auth.openai.com,
# and Druks polls for the grant.
redirect_uri = "https://auth.openai.com/deviceauth/callback"
# This additional limit is a reserve quota, not a model's own quota.
reserve_limit_name = "gpt-reserve"

Expand Down Expand Up @@ -968,21 +983,65 @@ def _grant_body(cls, refresh_token: str) -> dict:
}

@classmethod
def authorize_url(cls, *, verifier: str, challenge: str) -> tuple[str, str]:
state = secrets.token_hex(16)
params = {
"id_token_add_organizations": "true",
"codex_cli_simplified_flow": "true",
"originator": "pi", # the only value verified against the live exchange
"client_id": cls._CLIENT_ID,
"response_type": "code",
"redirect_uri": cls.redirect_uri,
"scope": "openid profile email offline_access",
"code_challenge": challenge,
"code_challenge_method": "S256",
"state": state,
async def start_connection(cls, *, account_id: str | None = None) -> dict:
device = await post_token(
url="https://auth.openai.com/api/accounts/deviceauth/usercode",
body={"client_id": cls._CLIENT_ID},
form=False,
)
flow_id = secrets.token_urlsafe(24)
pending = json.dumps(
{
"device_auth_id": device["device_auth_id"],
"user_code": device["user_code"],
"account_id": account_id,
}
)
await get_client().set(
name=f"{CONNECT_PENDING_PREFIX}{flow_id}", value=pending, ex=_DEVICE_CODE_TTL_SECONDS
)
return {
"method": "device",
"connection_id": flow_id,
"authorize_url": "https://auth.openai.com/codex/device",
"user_code": device["user_code"],
"poll_interval": int(device["interval"]),
}
return f"https://auth.openai.com/oauth/authorize?{urlencode(params)}", state

@classmethod
async def check_connection(cls, *, flow_id: str) -> CompletedConnect | None:
key = f"{CONNECT_PENDING_PREFIX}{flow_id}"
pending = await get_client().get(key)
if not pending:
raise exceptions.ConnectError("This connect attempt expired. Start it again.")
device = json.loads(pending)

try:
async with httpx.AsyncClient(timeout=_TOKEN_REQUEST_TIMEOUT_SECONDS) as client:
response = await client.post(
"https://auth.openai.com/api/accounts/deviceauth/token",
json={
"device_auth_id": device["device_auth_id"],
"user_code": device["user_code"],
},
)
except httpx.HTTPError as error:
raise exceptions.ConnectError("The request to OpenAI failed. Try again.") from error
# OpenAI answers 403 or 404 until the operator approves.
if response.status_code in (403, 404):
return
if response.status_code != 200:
raise exceptions.ConnectError(
f"OpenAI rejected the device code (HTTP {response.status_code}). Try again."
)

await get_client().delete(key)
approval = response.json()
return await cls._complete(
code=approval["authorization_code"],
verifier=approval["code_verifier"],
account_id=device["account_id"],
)

@classmethod
async def exchange(cls, *, code: str, verifier: str) -> tuple[dict, str | None]:
Expand Down
55 changes: 49 additions & 6 deletions backend/druks/harnesses/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from contextlib import suppress

from fastapi import APIRouter, Body, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession

from druks.accounts.dependencies import current_session_account, current_session_or_setup
from druks.accounts.models import Account
Expand All @@ -12,10 +13,12 @@
from druks.secrets.models import VaultSecret

from . import directory
from .datastructures import CompletedConnect
from .exceptions import ConnectError
from .models import ProviderCatalog
from .providers import Provider, get_provider, get_providers, is_registered
from .schemas import (
ConnectChallengeResponse,
ProviderCatalogResponse,
ProviderDirectoryResponse,
ProviderKeyResponse,
Expand Down Expand Up @@ -89,14 +92,21 @@ def _resolve_provider(provider_id: str) -> Provider:
raise HTTPException(status_code=404, detail=f"Unknown provider: {provider_id!r}") from error


@router.post("/{provider_id}/connection/start")
@router.post(
"/{provider_id}/connection/start",
response_model=ConnectChallengeResponse,
response_model_by_alias=True,
response_model_exclude_none=True,
)
async def start_connection(
provider_id: str, account: Account | None = Depends(current_session_or_setup)
) -> dict[str, str]:
) -> dict:
provider = _resolve_provider(provider_id)
# In none/zero the flow starts unbound, and its completion creates the operator.
url, flow_id = await provider.connect_start(account_id=account.id if account else None)
return {"authorizeUrl": url, "connectionId": flow_id}
try:
# In none/zero the flow starts unbound, and its completion creates the operator.
return await provider.start_connection(account_id=account.id if account else None)
except ConnectError as error:
raise HTTPException(status_code=422, detail=str(error)) from error


@router.post(
Expand All @@ -113,9 +123,42 @@ async def complete_connection(
) -> AccountResponse:
provider = _resolve_provider(provider_id)
try:
completed = await provider.connect_complete(flow_id=flow_id, pasted=code)
completed = await provider.complete_connection(flow_id=flow_id, pasted=code)
except ConnectError as exc:
raise HTTPException(status_code=422, detail=str(exc)) from exc
return await _store_connection(session, provider=provider, completed=completed, account=account)


@router.post(
"/{provider_id}/connection/check",
response_model=AccountResponse | None,
response_model_by_alias=True,
)
async def check_connection(
session: SessionDep,
provider_id: str,
account: Account | None = Depends(current_session_or_setup),
flow_id: str = Body(..., embed=True, alias="connectionId"),
) -> AccountResponse | None:
"""The connected account after the operator approves. None until then."""
provider = _resolve_provider(provider_id)
try:
completed = await provider.check_connection(flow_id=flow_id)
except ConnectError as error:
raise HTTPException(status_code=422, detail=str(error)) from error
if completed:
return await _store_connection(
session, provider=provider, completed=completed, account=account
)


async def _store_connection(
session: AsyncSession,
*,
provider: Provider,
completed: CompletedConnect,
account: Account | None,
) -> AccountResponse:
if account and account.id == completed.account_id:
resolved = account
elif completed.account_id:
Expand Down
10 changes: 9 additions & 1 deletion backend/druks/harnesses/schemas.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from datetime import datetime
from typing import Annotated
from typing import Annotated, Literal

from pydantic import AliasPath, BeforeValidator, ConfigDict, Field

Expand All @@ -17,6 +17,14 @@ class ProviderResponse(Schema):
billing_options: SortedNames


class ConnectChallengeResponse(Schema):
method: Literal["code", "device"]
connection_id: str
authorize_url: str
user_code: str | None = None
poll_interval: int | None = None


class ProviderSubscriptionResponse(Schema):
model_config = ConfigDict(from_attributes=True)

Expand Down
2 changes: 2 additions & 0 deletions backend/tests/test_auth_boundary.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
"/api/providers",
"/api/providers/{provider_id}/connection/start",
"/api/providers/{provider_id}/connection/complete",
"/api/providers/{provider_id}/connection/check",
"/api/secrets/{identity_id}/{name}", # a box's identity bearer, nothing else
"/api/{path:path}", # the JSON-404 catch-all
}
Expand Down Expand Up @@ -105,6 +106,7 @@
"/api/providers",
"/api/providers/{provider_id}/connection/start",
"/api/providers/{provider_id}/connection/complete",
"/api/providers/{provider_id}/connection/check",
}


Expand Down
37 changes: 28 additions & 9 deletions backend/tests/test_identity.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,12 +49,18 @@ def _connect(
email: str = "me@example.com",
headers: dict[str, str] | None = None,
):
if provider == "openai":
_mock_exchange_codex(monkeypatch, email=email)
start = client.post(f"/api/providers/{provider}/connection/start", headers=headers)
assert start.status_code == 200
if provider == "anthropic":
_mock_exchange(monkeypatch, _grant(email))
else:
_mock_exchange_codex(monkeypatch, email=email)
return client.post(
f"/api/providers/{provider}/connection/check",
json={"connectionId": start.json()["connectionId"]},
headers=headers,
)
return client.post(
f"/api/providers/{provider}/connection/complete",
json={"code": "thecode", "connectionId": start.json()["connectionId"]},
Expand All @@ -70,10 +76,21 @@ def _mock_exchange_codex(monkeypatch, *, email: str):
}
header = base64.urlsafe_b64encode(b'{"alg":"none"}').rstrip(b"=").decode()
payload = base64.urlsafe_b64encode(json.dumps(claims).encode()).rstrip(b"=").decode()
_mock_exchange(
monkeypatch,
{"access_token": f"{header}.{payload}.sig", "refresh_token": "RT", "id_token": "ID"},
)

async def fake_post(self, url, **_kwargs):
if url.endswith("/usercode"):
grant = {"device_auth_id": "device-id", "user_code": "ABCD-EFGH", "interval": "1"}
elif url.endswith("/deviceauth/token"):
grant = {"authorization_code": "approved", "code_verifier": "verifier"}
else:
grant = {
"access_token": f"{header}.{payload}.sig",
"refresh_token": "RT",
"id_token": "ID",
}
return httpx.Response(200, json=grant, request=httpx.Request("POST", url))

monkeypatch.setattr(providers.httpx.AsyncClient, "post", fake_post)


async def test_header_mode_requires_exactly_one_nonblank_assertion(tmp_path, druks_db):
Expand Down Expand Up @@ -292,6 +309,7 @@ async def test_concurrent_setup_completions_with_one_email_converge(
with _client(tmp_path) as client:
# No account exists yet, so both flows start unbound.
first = client.post("/api/providers/anthropic/connection/start")
_mock_exchange_codex(monkeypatch, email="me@example.com")
second = client.post("/api/providers/openai/connection/start")
_mock_exchange(monkeypatch, _grant("me@example.com"))
assert (
Expand All @@ -304,8 +322,8 @@ async def test_concurrent_setup_completions_with_one_email_converge(
_mock_exchange_codex(monkeypatch, email="me@example.com")
assert (
client.post(
"/api/providers/openai/connection/complete",
json={"code": "c2", "connectionId": second.json()["connectionId"]},
"/api/providers/openai/connection/check",
json={"connectionId": second.json()["connectionId"]},
).status_code
== 200
)
Expand All @@ -318,6 +336,7 @@ async def test_a_stale_unbound_completion_attaches_to_the_operator(tmp_path, mon
# The first completion creates the operator. The second, with another
# provider email, attaches to it instead of creating a second account.
first = client.post("/api/providers/anthropic/connection/start")
_mock_exchange_codex(monkeypatch, email="me@example.com")
second = client.post("/api/providers/openai/connection/start")
_mock_exchange(monkeypatch, _grant("a@example.com"))
client.post(
Expand All @@ -326,8 +345,8 @@ async def test_a_stale_unbound_completion_attaches_to_the_operator(tmp_path, mon
)
_mock_exchange_codex(monkeypatch, email="b@example.com")
completed = client.post(
"/api/providers/openai/connection/complete",
json={"code": "c2", "connectionId": second.json()["connectionId"]},
"/api/providers/openai/connection/check",
json={"connectionId": second.json()["connectionId"]},
)
assert completed.status_code == 200
assert completed.json()["username"] == "a@example.com"
Expand Down
Loading
Loading