Skip to content
Open
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
47 changes: 45 additions & 2 deletions coworker/overrides.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,14 +34,36 @@ def _specificity(pattern: str) -> int:


class RiskOverrideStore:
"""Rules live in one JSON file; every instance watches its mtime and reloads
lazily, so a REST write through one instance is seen by the per-engine
instances already captured in live PermissionEngines — no rebuild needed."""

def __init__(self, path: Optional[str | Path] = None) -> None:
self.path = Path(path) if path else None
self._rules: list[_Rule] = self._load()
self._mtime: Optional[float] = None
self._rules: list[_Rule] = []
self._refresh()

def _refresh(self) -> None:
"""Reload from disk iff the file changed since we last read it."""
if not self.path:
return
try:
mtime = self.path.stat().st_mtime
except OSError:
self._rules, self._mtime = [], None
return
if mtime == self._mtime:
return
self._rules, self._mtime = self._load(), mtime

def _load(self) -> list[_Rule]:
if not (self.path and self.path.is_file()):
return []
data = json.loads(self.path.read_text(encoding="utf-8"))
try:
data = json.loads(self.path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return []
rules = []
for r in data.get("rules", []):
try:
Expand All @@ -66,15 +88,36 @@ def save(self) -> None:
),
encoding="utf-8",
)
try:
self._mtime = self.path.stat().st_mtime
except OSError:
self._mtime = None

def set_rule(self, pattern: str, risk: RiskClass | str) -> None:
"""Add/replace a user override (the everyday path writes this from the approval UI)."""
risk = RiskClass(risk) if not isinstance(risk, RiskClass) else risk
self._refresh()
self._rules = [r for r in self._rules if r.pattern != pattern]
self._rules.append(_Rule(pattern, risk))
self.save()

def remove_rule(self, pattern: str) -> bool:
"""Delete an override by exact pattern. True if a rule was removed."""
self._refresh()
before = len(self._rules)
self._rules = [r for r in self._rules if r.pattern != pattern]
if len(self._rules) == before:
return False
self.save()
return True

def rules(self) -> list[dict]:
"""Current rules as plain dicts (REST/GUI listing)."""
self._refresh()
return [{"pattern": r.pattern, "risk": r.risk.value} for r in self._rules]

def resolve(self, tool_name: str) -> Optional[RiskClass]:
self._refresh()
best: Optional[RiskClass] = None
best_score = -1
for r in self._rules:
Expand Down
14 changes: 14 additions & 0 deletions coworker/server/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -793,6 +793,20 @@ async def mcp_oauth_callback(
)
)

@app.get("/v1/risk-overrides")
def risk_overrides_list() -> dict[str, Any]:
return manager.list_risk_overrides()

@app.post("/v1/risk-overrides")
def risk_overrides_set(body: dict) -> dict[str, Any]:
# User-local trust decision (GUI/REST). Personas/packages never reach this
# path — the no-self-grant rule lives in the store's design, not here.
return manager.set_risk_override(body or {})

@app.delete("/v1/risk-overrides")
def risk_overrides_delete(pattern: str) -> dict[str, Any]:
return manager.delete_risk_override(pattern)

@app.post("/v1/mcp/reload")
async def mcp_reload() -> dict[str, Any]:
return await manager.reload_mcp()
Expand Down
52 changes: 50 additions & 2 deletions coworker/server/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,12 @@ def __init__(
self.secrets, default_provider="openai", on_use=self._note_provider_use
)
self.mcp = MCPManager(secrets=self.secrets)
# The user-local risk-override store (Phase 2). Same file as the per-engine
# instances in build_engine; the store's mtime watch makes REST writes here
# visible to live engines without a rebuild.
from ..overrides import RiskOverrideStore

self.risk_overrides = RiskOverrideStore(state_dir() / "risk_overrides.json")
# OAuth MCP servers with a sign-in in flight / their last connect error —
# feeds list_mcp's status so the GUI can show "authorizing…" and failures.
self._mcp_authorizing: set[str] = set()
Expand Down Expand Up @@ -1125,12 +1131,54 @@ async def mcp_tools(self, name: str) -> dict[str, Any]:
"name": name,
"ok": True,
"tools": [
{"name": t.name, "description": getattr(t, "description", "")}
for t in conn.tools
self._mcp_tool_row(server, t) for t in conn.tools
],
}
return {"name": name, "ok": False, "error": "unknown server", "tools": []}

def _mcp_tool_row(self, server: Any, t: Any) -> dict[str, Any]:
"""One tool listing row, enriched with its registry name and effective risk
so the GUI can render per-tool risk controls (Phase 2 overrides)."""
from types import SimpleNamespace

from ..mcp.tools import tool_name as mcp_tool_name
from ..risk import classify

full = mcp_tool_name(server.name, t.name)
meta = SimpleNamespace(requires_approval=server.requires_approval, category="mcp")
return {
"name": t.name,
"description": getattr(t, "description", ""),
"full_name": full,
"risk": classify(full, meta, self.risk_overrides.resolver()).value,
"default_risk": classify(full, meta).value,
"overridden": self.risk_overrides.resolve(full) is not None,
}

# ---- risk overrides (Phase 2 surface) -----------------------------------
def list_risk_overrides(self) -> dict[str, Any]:
return {"rules": self.risk_overrides.rules()}

def set_risk_override(self, payload: dict[str, Any]) -> dict[str, Any]:
from ..risk import RiskClass

pattern = str(payload.get("pattern") or "").strip()
if not pattern:
return {"ok": False, "error": "pattern is required"}
try:
risk = RiskClass(str(payload.get("risk") or ""))
except ValueError:
return {
"ok": False,
"error": f"risk must be one of {[r.value for r in RiskClass]}",
}
self.risk_overrides.set_rule(pattern, risk)
return {"ok": True, "rules": self.risk_overrides.rules()}

def delete_risk_override(self, pattern: str) -> dict[str, Any]:
removed = self.risk_overrides.remove_rule(pattern)
return {"ok": removed, "rules": self.risk_overrides.rules()}

async def reload_mcp(self) -> dict[str, Any]:
"""Drop live MCP connections so new sessions reconnect with fresh config."""
await self.mcp.aclose()
Expand Down
37 changes: 37 additions & 0 deletions surfaces/gui/e2e/fixtures.ts
Original file line number Diff line number Diff line change
Expand Up @@ -548,6 +548,8 @@ export async function mockApi(page: import("@playwright/test").Page) {
const automations: any[] = [{ ...AUTOMATION }, { ...AUTOMATION_CLEAN }];
// MCP servers (empty by default; the granola OAuth quick-add test populates it).
const mcpServers: any[] = [];
// User-local risk overrides (Phase 2): pattern → risk, mutated by /v1/risk-overrides.
const riskOverrides: { pattern: string; risk: string }[] = [];
const automationRuns: any[] = AUTOMATION_RUNS.map((r) => ({ ...r }));
// Per-session unattended flag — mutable so the composer's "Send to Inbox" toggle persists and
// the app reads it back (which is what gates parking approvals to the Inbox vs an inline card).
Expand Down Expand Up @@ -1562,6 +1564,41 @@ export async function mockApi(page: import("@playwright/test").Page) {
}
return json({ ok: true });
}
// Tool listing with per-tool effective risk (Phase 2 overrides). Two tools:
// by MCP default both are external ("asks") until a user override relaxes one.
const mt = p.match(/\/v1\/mcp\/([^/]+)\/tools$/);
if (mt && m === "GET") {
const server = decodeURIComponent(mt[1]);
const tools = ["get_status", "set_value"].map((name) => {
const full = `mcp__${server}__${name}`;
const ov = riskOverrides.find((r) => r.pattern === full);
return {
name,
description: `${name} on ${server}`,
full_name: full,
risk: ov ? ov.risk : "external",
default_risk: "external",
overridden: !!ov,
};
});
return json({ name: server, ok: true, tools });
}
}
if (p.endsWith("/v1/risk-overrides")) {
if (m === "GET") return json({ rules: riskOverrides });
if (m === "POST") {
const b = req.postDataJSON();
const i = riskOverrides.findIndex((r) => r.pattern === b.pattern);
if (i >= 0) riskOverrides.splice(i, 1);
riskOverrides.push({ pattern: b.pattern, risk: b.risk });
return json({ ok: true, rules: riskOverrides });
}
if (m === "DELETE") {
const pat = new URL(req.url()).searchParams.get("pattern");
const i = riskOverrides.findIndex((r) => r.pattern === pat);
if (i >= 0) riskOverrides.splice(i, 1);
return json({ ok: i >= 0, rules: riskOverrides });
}
}
if (p.endsWith("/v1/unrouted")) return json([]);

Expand Down
37 changes: 37 additions & 0 deletions surfaces/gui/e2e/mcp-tool-risk.spec.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
// Per-tool MCP risk overrides (Phase 2): the tools list shows each tool's
// effective risk, and clicking a tool toggles a user-local "trust as read-only"
// override — relaxed tools run without approval prompts, and the override is
// removable in place.
import { expect } from "@playwright/test";
import { test } from "./fixtures";

async function openMcpTab(page) {
await page.goto("/");
await page.getByTestId("account-row").click();
await page.getByRole("button", { name: "Connectors", exact: true }).click();
await page.getByRole("button", { name: "MCP servers", exact: true }).click();
}

test("tool chips show risk and toggle a read override", async ({ page }) => {
await openMcpTab(page);

// Connect the curated Granola preset so a server row exists (mock flips it).
await page.getByTestId("mcp-preset-granola").getByRole("button", { name: "Connect" }).click();
const row = page.locator(".space-y-2 > div").filter({ hasText: "granola" }).first();
await expect(row).toContainText("connected", { timeout: 10_000 });

// Open the tools list: both tools ask by MCP's conservative default.
await row.getByRole("button", { name: "tools", exact: true }).click();
const statusChip = row.getByTestId("mcp-tool-risk-get_status");
await expect(statusChip).toContainText("asks");
await expect(row.getByTestId("mcp-tool-risk-set_value")).toContainText("asks");

// Trust the read tool: chip flips to auto (runs without asking); the other stays gated.
await statusChip.click();
await expect(statusChip).toContainText("auto");
await expect(row.getByTestId("mcp-tool-risk-set_value")).toContainText("asks");

// Click again: the override is removed and approval is restored.
await statusChip.click();
await expect(statusChip).toContainText("asks");
});
33 changes: 32 additions & 1 deletion surfaces/gui/src/api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -306,13 +306,44 @@ export async function deleteMcpServer(name: string) {
return res.json();
}

export interface McpToolRow {
name: string;
description: string;
/** Registry name (`mcp__<server>__<tool>`) — the pattern risk overrides match. */
full_name: string;
/** Effective risk with the user's overrides applied. */
risk: string;
/** What the tool would be without overrides (MCP default: external → asks). */
default_risk: string;
overridden: boolean;
}

export async function getMcpTools(
name: string,
): Promise<{ ok: boolean; error?: string; tools: { name: string; description: string }[] }> {
): Promise<{ ok: boolean; error?: string; tools: McpToolRow[] }> {
const res = await fetch(`${httpBase()}/v1/mcp/${encodeURIComponent(name)}/tools`);
return res.json();
}

// User-local risk overrides (Phase 2): relax a trusted MCP tool to `read` so it
// stops prompting, or remove the override to restore the conservative default.
export async function setRiskOverride(pattern: string, risk: string) {
const res = await fetch(`${httpBase()}/v1/risk-overrides`, {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({ pattern, risk }),
});
return res.json();
}

export async function deleteRiskOverride(pattern: string) {
const res = await fetch(
`${httpBase()}/v1/risk-overrides?pattern=${encodeURIComponent(pattern)}`,
{ method: "DELETE" },
);
return res.json();
}

export async function reloadMcp() {
const res = await fetch(`${httpBase()}/v1/mcp/reload`, { method: "POST" });
return res.json();
Expand Down
57 changes: 46 additions & 11 deletions surfaces/gui/src/components/ManageTabs.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,9 @@ import {
disallowUser,
getMcpServers,
getMcpTools,
type McpToolRow,
setRiskOverride,
deleteRiskOverride,
signoutMcp,
getSettings,
getSubscriptions,
Expand Down Expand Up @@ -360,7 +363,7 @@ function McpRow({
onRemove: () => void;
onRefresh: () => void;
}) {
const [tools, setTools] = useState<{ name: string; description: string }[] | null>(null);
const [tools, setTools] = useState<McpToolRow[] | null>(null);
const [busy, setBusy] = useState(false);
const [toolErr, setToolErr] = useState<string | null>(null);

Expand Down Expand Up @@ -433,17 +436,49 @@ function McpRow({
)}
{toolErr && <div className="text-[12.5px] text-danger mt-1.5">{toolErr}</div>}
{tools && (
<div className="mt-2.5 pt-2.5 border-t border-line flex flex-wrap gap-1.5">
<div className="mt-2.5 pt-2.5 border-t border-line">
{tools.length === 0 && <div className="text-[12px] text-faint">No tools.</div>}
{tools.map((t) => (
<span
key={t.name}
title={t.description}
className="font-mono text-[11.5px] px-1.5 py-0.5 rounded-md bg-paper border border-line"
>
{t.name}
</span>
))}
{tools.length > 0 && (
<div className="text-[11.5px] text-faint mb-1.5">
Click a tool to let it run without asking (marks it read-only for
this machine); click again to restore approval.
</div>
)}
<div className="flex flex-wrap gap-1.5">
{tools.map((t) => {
const relaxed = t.overridden && t.risk === "read";
const asks = t.risk !== "read";
return (
<button
key={t.name}
data-testid={`mcp-tool-risk-${t.name}`}
title={
t.description +
(relaxed
? "\n\nOverride active: runs without asking. Click to restore approval."
: asks
? "\n\nAsks for approval before each call. Click to trust as read-only."
: "\n\nRead-only: runs without asking.")
}
onClick={async () => {
if (t.overridden) await deleteRiskOverride(t.full_name);
else await setRiskOverride(t.full_name, "read");
const res = await getMcpTools(server.name);
if (res.ok) setTools(res.tools);
}}
className={
"font-mono text-[11.5px] px-1.5 py-0.5 rounded-md bg-paper border " +
(relaxed ? "border-ink/60 text-ink" : "border-line text-muted hover:text-ink")
}
>
{t.name}
<span className={"ml-1.5 " + (relaxed ? "text-ink" : "text-faint")}>
{relaxed ? "auto" : asks ? "asks" : "read"}
</span>
</button>
);
})}
</div>
</div>
)}
</div>
Expand Down
Loading