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
5 changes: 5 additions & 0 deletions backend/druks/agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
from druks.sandbox.client import provisioning_key, sandbox_client
from druks.sandbox.models import SandboxIdentity, SecretRef
from druks.sandbox.templates import get_template_id
from druks.secrets.enums import SecretKind
from druks.settings import load_settings
from druks.usage.models import UsageScrape
from druks.workflows import _in_step, current_workflow
Expand Down Expand Up @@ -172,6 +173,10 @@ def __post_init__(self) -> None:
f"{declared} goes to a sandbox, so {secret.service.__name__} must declare "
"`host`: the one host the secret may be sent to"
)
if secret.service.secret_kind != SecretKind.STATIC:
raise TypeError(
f"{declared} belongs to an App key, which issues tokens; no sandbox holds it"
)
if self.id: # an explicit id means a standalone agent — it registers itself now
agents.register(self)

Expand Down
29 changes: 22 additions & 7 deletions backend/druks/chat/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

from druks.accounts.enums import AccountKind
from druks.apps.loader import get_app
from druks.apps.registry import channels
from druks.apps.registry import channels, services
from druks.durable.engine import step_session
from druks.durable.models import Run
from druks.files.constants import MAX_UPLOAD_BYTES
Expand All @@ -39,6 +39,8 @@
from druks.sandbox.layout import get_remote_home, get_work_root
from druks.sandbox.models import SandboxIdentity, SecretRef
from druks.sandbox.templates import get_template_id
from druks.secrets.datastructures import Audience
from druks.secrets.models import VaultSecret
from druks.services.exceptions import ServiceNotConnectedError
from druks.workspaces import Workspace

Expand Down Expand Up @@ -107,15 +109,15 @@ async def get_sandbox(
*,
config: AgentConfig,
allowed_tools: AllowedTools,
mcp_secret_refs: list[SecretRef],
secret_refs: list[SecretRef],
) -> tuple[Host, SandboxIdentity]:
"""The account's sandbox for a new turn: a live sandbox that holds the Chat
agent's current secrets, or a new one."""
server = get_druks_mcp_server(allowed_tools=())
token = await get_druks_account_token(session, account_id, allowed_tools, name=CHAT_KEY_NAME)
refs = [
*config.secret_refs,
*mcp_secret_refs,
*secret_refs,
SecretRef(
name=get_bearer_token_env_var(server.name).lower(),
secret_id=token.id,
Expand Down Expand Up @@ -208,19 +210,32 @@ async def deliver_pending(session: AsyncSession, conversation: Conversation) ->
f"Chat runs on {adapters}. The Bot's settings select "
f"{config.harness_class.name}. Set its harness to one of them."
)
# A bot serves outside people, so only the operator reaches the enabled MCP servers.
mcp_servers, mcp_secret_refs = (), []
# A bot serves outside people, so only the operator reaches the enabled MCP
# servers and holds their own sign-ins.
mcp_servers, secret_refs = (), []
if conversation.account.kind == AccountKind.OPERATOR:
# A server the operator has not connected must not stop the chat.
mcp_servers, mcp_secret_refs = await Workspace.get_all_mcp_servers(
mcp_servers, secret_refs = await Workspace.get_all_mcp_servers(
session, None, conversation.account_id, skip_unauthenticated=True
)
for service in services.all():
if service.host and service.token_endpoint:
sign_ins = await VaultSecret.list_account_connections(
session, Audience.service(service.slug), conversation.account_id
)
# A sandbox has one variable per service, so it holds one sign-in.
secret_refs += [
SecretRef(
name=service.slug, secret_id=sign_in.id, host=service.host
)
for sign_in in sign_ins[:1]
]
host, identity = await get_sandbox(
session,
conversation.account_id,
config=config,
allowed_tools=tools,
mcp_secret_refs=mcp_secret_refs,
secret_refs=secret_refs,
)
try:
bridge = Bridge(host)
Expand Down
3 changes: 2 additions & 1 deletion backend/druks/contrib/software_factory/workflows.py
Original file line number Diff line number Diff line change
Expand Up @@ -484,9 +484,10 @@ async def get_secrets(cls, subject: Any) -> list[SandboxSecret]:
actor = await get_review_actor()
return [
SandboxSecret(
name=Github.secret_name,
name=actor.service.slug,
secret_id=(await actor.service.get()).id,
resource=cls.get_repo(subject),
host=actor.service.host,
)
]

Expand Down
9 changes: 7 additions & 2 deletions backend/druks/core/services.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,8 @@ class Github(Service):
person's sign-in through it links their GitHub account to their Druks account."""

secret_kind = SecretKind.APP_KEY
# The Drukbox catalog name a box holds this identity's token under.
secret_name = "github"
# Drukbox knows this host: a token for it reaches git and gh.
host = "github.com"
description = (
"The GitHub App druks acts as. Create it from here, or paste an existing "
"App's credentials from the GitHub developer settings page."
Expand Down Expand Up @@ -105,6 +105,11 @@ async def issue_token(cls, resource: str) -> tuple[str, datetime]:
"""The installation token for the repo, and the expiry GitHub gave it."""
return await (await cls.get_client()).token_for_repo(resource)

@classmethod
def is_grant_revoked(cls, status: int, tokens: dict[str, Any]) -> bool:
# GitHub answers a dead refresh token with a 200.
return tokens.get("error") == "bad_refresh_token"

@classmethod
async def get_identity(cls, access_token: str) -> dict[str, Any]:
"""The person behind a user token, keyed the way ``Account.lookup`` finds them."""
Expand Down
3 changes: 2 additions & 1 deletion backend/druks/sandbox/datastructures.py
Original file line number Diff line number Diff line change
Expand Up @@ -162,7 +162,8 @@ class SandboxSecret:
"""A secret a workspace's box holds as a placeholder. ``secret_id`` names the
vault row the issuer answers from and ``resource`` what its token is for. A
``host`` makes it a custom entry: the proxy swaps the placeholder in the
request header at that host, and the box reads it from ``name.upper()``."""
request header at that host, and the box reads it from ``name.upper()``.
Drukbox knows ``github.com``: that entry is its GitHub service."""

name: str
secret_id: str
Expand Down
3 changes: 2 additions & 1 deletion backend/druks/sandbox/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,8 @@ class SecretRef(Base):
# What the token is for: the repo. Empty for a subscription.
resource: Mapped[str] = mapped_column(default="")
# The host the box's placeholder is swapped at, for a custom entry such
# as an MCP server. Empty for a Drukbox catalog entry.
# as an MCP server, or ``github.com`` for Drukbox's GitHub service. Empty
# for a Drukbox catalog entry.
host: Mapped[str] = mapped_column(default="")

identity: Mapped["SandboxIdentity"] = relationship(back_populates="secret_refs", lazy="raise")
Expand Down
9 changes: 8 additions & 1 deletion backend/druks/services/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
from druks.secrets.models import VaultSecret

from .exceptions import OauthExchangeError, ServiceConnectError, ServiceNotConnectedError
from .oauth import OauthClient, fetch_identity
from .oauth import OauthClient, fetch_identity, is_grant_revoked

# GoogleCalendar -> google_calendar, HTTPServer -> http_server.
_CAMEL_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])|(?<=[A-Z])(?=[A-Z][a-z])")
Expand Down Expand Up @@ -317,6 +317,12 @@ def read_grant(cls, tokens: dict[str, Any]) -> dict[str, Any]:
"scopes": tokens.get("scope", "").split(),
}

@classmethod
def is_grant_revoked(cls, status: int, tokens: dict[str, Any]) -> bool:
"""Whether the token endpoint's answer to a refresh says the provider revoked
the grant. Override for a provider that reports it otherwise than RFC 6749."""
return is_grant_revoked(status, tokens)

@classmethod
async def get_oauth_client(cls) -> OauthClient:
"""The connected identity as a configured ``OauthClient``, keyed by
Expand All @@ -333,6 +339,7 @@ async def get_oauth_client(cls) -> OauthClient:
client_secret=connected.secrets["client_secret"],
basic_auth=cls.basic_auth,
extra_authorize_params=cls.extra_authorize_params,
is_grant_revoked=cls.is_grant_revoked,
)

@classmethod
Expand Down
49 changes: 31 additions & 18 deletions backend/druks/services/oauth.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import hashlib
import json
import secrets
from collections.abc import Callable
from contextlib import AsyncExitStack
from datetime import UTC, datetime, timedelta
from typing import Any, cast
Expand Down Expand Up @@ -32,6 +33,12 @@ def _http() -> httpx.AsyncClient:
return httpx.AsyncClient(timeout=30.0, follow_redirects=True)


def is_grant_revoked(status: int, tokens: dict[str, Any]) -> bool:
"""Whether a token endpoint's answer says the provider revoked the grant, as RFC
6749 reports it. A service overrides its own method for a provider that differs."""
return status != 200 and tokens.get("error") == "invalid_grant"


async def _post_token(
http: httpx.AsyncClient,
token_endpoint: str,
Expand Down Expand Up @@ -84,6 +91,7 @@ class OauthClient:
example RFC 8707's ``resource``. ``extra_authorize_params`` go into every
consent query, for example Google's ``access_type=offline`` and
``prompt=consent``. Each ``begin_connect`` asks for its own scopes.
``is_grant_revoked`` reads the token endpoint's answer to a refresh.
"""

def __init__(
Expand All @@ -99,6 +107,7 @@ def __init__(
extra_authorize_params: dict[str, str] | None = None,
mint_wait_interval_seconds: float = OAUTH_MINT_WAIT_INTERVAL_SECONDS,
mint_wait_attempts: int = OAUTH_MINT_WAIT_ATTEMPTS,
is_grant_revoked: Callable[[int, dict[str, Any]], bool] = is_grant_revoked,
) -> None:
self.provider = provider
self.authorization_endpoint = authorization_endpoint
Expand All @@ -110,6 +119,7 @@ def __init__(
self.extra_authorize_params = dict(extra_authorize_params or {})
self.mint_wait_interval_seconds = mint_wait_interval_seconds
self.mint_wait_attempts = mint_wait_attempts
self.is_grant_revoked = is_grant_revoked

async def begin_connect(
self,
Expand Down Expand Up @@ -241,30 +251,33 @@ async def get_access_token(
)
except httpx.HTTPError as error:
raise OauthRefreshError(self.provider, str(error)) from error
if response.status_code != 200:
if "invalid_grant" in response.text:
# The provider withdrew the grant; presenting it again can never
# succeed. The revoke commits on its own: the caller's step
# session rolls back when this error propagates.
async with get_session(session.bind) as own:
revoked = await own.get(VaultSecret, connection.id)
await self.disconnect(revoked, reason="invalid_grant")
await own.commit()
try:
tokens = response.json()
except ValueError:
tokens = None
if not isinstance(tokens, dict):
if response.status_code == 200:
raise OauthRefreshError(
self.provider,
"the provider revoked the grant; sign in again to restore the connection",
self.provider, "the token endpoint returned malformed JSON"
)
tokens = {}
if self.is_grant_revoked(response.status_code, tokens):
# Presenting the grant again can never succeed. The revoke commits on
# its own: the caller's step session rolls back when this error propagates.
async with get_session(session.bind) as own:
revoked = await own.get(VaultSecret, connection.id)
await self.disconnect(revoked, reason="invalid_grant")
await own.commit()
raise OauthRefreshError(
self.provider,
"the provider revoked the grant; sign in again to restore the connection",
)
if response.status_code != 200:
await redis.delete(token_key)
raise OauthRefreshError(
self.provider, f"HTTP {response.status_code} from the token endpoint"
)
try:
tokens = response.json()
except ValueError as error:
raise OauthRefreshError(
self.provider, "the token endpoint returned malformed JSON"
) from error
if not isinstance(tokens, dict) or not tokens.get("access_token"):
if not tokens.get("access_token"):
raise OauthRefreshError(
self.provider, "the token endpoint returned no access token"
)
Expand Down
3 changes: 2 additions & 1 deletion backend/druks/workspaces.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,9 +278,10 @@ async def get_secrets(cls, subject: Any) -> list[SandboxSecret]:
# issuer reads. A service that is not connected fails here, before the box.
return [
SandboxSecret(
name=cls.github.secret_name,
name=cls.github.slug,
secret_id=(await cls.github.get()).id,
resource=cls.get_repo(subject),
host=cls.github.host,
)
]

Expand Down
4 changes: 3 additions & 1 deletion backend/tests/druks-field_notes/tests/test_workflows.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,4 +48,6 @@ async def test_survey_workspace_clones_the_subject_repo(druks_db):
"github", identity={"app_id": "1", "slug": "druks-operator"}, secrets={"private_key": "pem"}
)
[secret] = await workflow.workspace_class.get_secrets(await workflow.subject)
assert secret == SandboxSecret(name="github", secret_id=row.id, resource="acme/widgets")
assert secret == SandboxSecret(
name="github", secret_id=row.id, resource="acme/widgets", host="github.com"
)
8 changes: 6 additions & 2 deletions backend/tests/software_factory/test_build_workspace.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,7 +85,9 @@ async def test_build_workspace_declares_its_github_mcp_as_the_review_actor(druks
assert github.bearer_token_env_var == get_bearer_token_env_var(GITHUB_MCP_NAME)
assert ref.key == ("mcp_github_token", reviewer.id, "o/main", "api.githubcopilot.com")
[clone] = await BuildWorkspace.get_secrets(subject)
assert clone == SandboxSecret(name="github", secret_id=operator.id, resource="o/main")
assert clone == SandboxSecret(
name="github", secret_id=operator.id, resource="o/main", host="github.com"
)


async def test_get_workspace_kwargs_carries_the_build_fields():
Expand Down Expand Up @@ -157,7 +159,9 @@ async def test_review_mcp_and_gh_use_the_review_actor(druks_db):
assert github.bearer_token_env_var == get_bearer_token_env_var(GITHUB_MCP_NAME)
assert ref.key == ("mcp_github_token", reviewer.id, "o/app", "api.githubcopilot.com")
[clone] = await ReviewWorkspace.get_secrets(subject)
assert clone == SandboxSecret(name="github", secret_id=reviewer.id, resource="o/app")
assert clone == SandboxSecret(
name="github_reviewer", secret_id=reviewer.id, resource="o/app", host="github.com"
)


class _IdentitySandbox:
Expand Down
2 changes: 1 addition & 1 deletion backend/tests/software_factory/test_review.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,7 +199,7 @@ async def test_an_unconnected_reviewer_borrows_the_operator_in_comment_mode(druk

def test_the_reviewer_is_an_optional_service_the_app_declares():
assert (GithubReviewer.slug, GithubReviewer.required) == ("github_reviewer", False)
assert GithubReviewer.secret_name == "github"
assert GithubReviewer.host == "github.com"
fields = GithubReviewer.Settings.model_fields
assert set(fields) == {"app_id", "private_key"}
assert field_kind(fields["private_key"]) == "secret"
Expand Down
24 changes: 17 additions & 7 deletions backend/tests/test_chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,14 +206,14 @@ async def attach(*, host_id):
conversation.account_id,
config=config,
allowed_tools=Toolkit.ALL,
mcp_secret_refs=[],
secret_refs=[],
)
second_host, _identity = await service.get_sandbox(
druks_db,
second.account_id,
config=config,
allowed_tools=Toolkit.ALL,
mcp_secret_refs=[],
secret_refs=[],
)
login = await get_druks_account_token(druks_db, conversation.account_id, (), name="login")
moved = SimpleNamespace(
Expand All @@ -224,7 +224,7 @@ async def attach(*, host_id):
second.account_id,
config=moved,
allowed_tools=Toolkit.ALL,
mcp_secret_refs=[],
secret_refs=[],
)

assert first_host.id == second_host.id
Expand Down Expand Up @@ -485,7 +485,7 @@ async def download(**values):
@pytest.mark.parametrize(
"kind, expected", [(AccountKind.OPERATOR, [LINEAR]), (AccountKind.BOT, [])]
)
async def test_only_the_operator_reaches_the_connected_mcp_servers(
async def test_only_the_operator_reaches_the_connected_mcp_servers_and_holds_their_sign_in(
druks_db, conversation, sandbox, monkeypatch, kind, expected
):
await McpServer.create(
Expand All @@ -495,6 +495,13 @@ async def test_only_the_operator_reaches_the_connected_mcp_servers(
secret_headers={"Authorization": "Bearer lin_secret"},
)
await McpServer.create(druks_db, name="sentry", url="https://mcp.sentry.dev/mcp", is_oauth=True)
sign_in = await VaultSecret.connect(
druks_db,
Audience.service("github"),
account_id=conversation.account_id,
refresh_token="ghr_one",
scopes=[],
)
conversation.account.kind = kind
starts = []

Expand All @@ -512,8 +519,11 @@ async def follow_turn(session, conversation, message, bridge):

await service.deliver_pending(druks_db, conversation)

[refs] = [call.kwargs["mcp_secret_refs"] for call in service.get_sandbox.await_args_list]
assert len(refs) == len(expected)
[refs] = [call.kwargs["secret_refs"] for call in service.get_sandbox.await_args_list]
assert [ref.name for ref in refs] == ["mcp_linear_header_0", "github"] * len(expected)
assert [ref.key for ref in refs if ref.name == "github"] == [
("github", sign_in.id, "", "github.com")
] * len(expected)
assert starts[0]["mcpServers"] == expected


Expand Down Expand Up @@ -939,7 +949,7 @@ async def test_stop_during_startup_keeps_the_sandbox_and_sends_only_the_next_mes
sandbox_requests = 0
prompts = []

async def get_sandbox(session, account_id, *, config, allowed_tools, mcp_secret_refs):
async def get_sandbox(session, account_id, *, config, allowed_tools, secret_refs):
nonlocal sandbox_requests
sandbox_requests += 1
await session.commit()
Expand Down
Loading
Loading