fix: tighten oauth and source resource handling
This commit is contained in:
@@ -6,12 +6,12 @@ param(
|
||||
[Parameter(ValueFromRemainingArguments = $true)]
|
||||
[string[]]$RemainingArgs
|
||||
)
|
||||
function funny([string] $output) {
|
||||
function New-PandocDiagramMetadata([string] $outputFormat) {
|
||||
return @{
|
||||
"diagram" = @{
|
||||
"engine" = @{
|
||||
"mermaid" = @{
|
||||
"outputFormat" = "$output"
|
||||
"outputFormat" = "$outputFormat"
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -29,7 +29,7 @@ else {
|
||||
exit 1
|
||||
}
|
||||
|
||||
$metadata = funny $outputFormat | ConvertTo-Json -Depth 10
|
||||
$metadata = New-PandocDiagramMetadata $outputFormat | ConvertTo-Json -Depth 10
|
||||
$pandoc_diagram = Join-Path $PSScriptRoot "../../stuff/pandoc-diagram.ps1"
|
||||
|
||||
$metatempfile = New-TemporaryFile
|
||||
|
||||
+18
-2
@@ -514,7 +514,8 @@ class OAuthCodeLoginFlow:
|
||||
provider: OAuthProviderConfig,
|
||||
client_id: str,
|
||||
client_secret: str | None,
|
||||
authorization_response: str,
|
||||
authorization_response: str | None,
|
||||
authorization_url_callback: Callable[[str, str], str | None] | None = None,
|
||||
) -> OAuthLoginResult:
|
||||
client = self._client_factory(
|
||||
client_id=client_id,
|
||||
@@ -522,7 +523,17 @@ class OAuthCodeLoginFlow:
|
||||
scope=" ".join(provider.scopes),
|
||||
code_challenge_method="S256",
|
||||
)
|
||||
client.create_authorization_url(str(provider.auth_url))
|
||||
authorization_url, state = client.create_authorization_url(
|
||||
str(provider.auth_url),
|
||||
redirect_uri=provider.redirect_uri,
|
||||
**provider.extra_authorize_params,
|
||||
)
|
||||
if authorization_url_callback is not None:
|
||||
callback_response = authorization_url_callback(authorization_url, state)
|
||||
if authorization_response is None:
|
||||
authorization_response = callback_response
|
||||
if authorization_response is None:
|
||||
raise ValueError("OAuth authorization response is required")
|
||||
token = await client.fetch_token(
|
||||
str(provider.token_url),
|
||||
authorization_response=authorization_response,
|
||||
@@ -536,6 +547,11 @@ class OAuthCodeLoginFlow:
|
||||
```
|
||||
|
||||
This helper supports pasted authorization response first. Browser callback can be a later refinement.
|
||||
`authorization_url_callback` is used by the CLI to show the generated URL and
|
||||
collect the pasted redirected callback URL; tests can inject it to avoid prompt
|
||||
I/O. Provider-specific authorization parameters such as Google's
|
||||
`access_type=offline` and `prompt=consent` belong in
|
||||
`provider.extra_authorize_params`, not in this generic helper.
|
||||
|
||||
- [ ] **Step 4: Run helper tests**
|
||||
|
||||
@@ -77,6 +77,40 @@ workflow-facing `CapabilitySource` objects. Provider-specific runtime pools,
|
||||
admin/apply hooks, auth, catalog caches, and live health checks stay outside
|
||||
this narrow seam until a source family needs them.
|
||||
|
||||
## Capability, Tool, Resource, And Prompt
|
||||
|
||||
Use these terms precisely:
|
||||
|
||||
- A **source** is the owner and namespace. Examples: `everything.default`,
|
||||
`wf.std`, `wf.source`, or a future OpenAPI/Python source.
|
||||
- A **tool** is provider-native. For MCP, it is an MCP tool discovered from an
|
||||
upstream server.
|
||||
- A **workflow capability** is workflow-native. It is the `NodeSpec` shape a
|
||||
graph can call with typed input, typed output, outcomes, validation, and trace
|
||||
behavior. Tools can be projected into workflow capabilities, but the two are
|
||||
not identical concepts.
|
||||
- A **resource** is source-owned addressable content. The URI is not globally
|
||||
meaningful by itself; it must be interpreted with the source that owns it.
|
||||
- A **prompt** is source-owned prompt/template inventory. Listing prompts is
|
||||
inventory; rendering a prompt is an upstream operation and may be stateful.
|
||||
|
||||
Saved workflow data should prefer logical source references over concrete
|
||||
source ids. For example, a resource ref should store:
|
||||
|
||||
```json
|
||||
{"logical_source": "drive", "uri": "gdrive://file/abc"}
|
||||
```
|
||||
|
||||
The deployment binding decides whether `drive` means `drive.personal`,
|
||||
`drive.work`, or another concrete source. Platform sources such as `wf.std` and
|
||||
`wf.source` are special because their logical source id is also their concrete
|
||||
source id, so they do not require deployment bindings.
|
||||
|
||||
Runtime dereference is explicit. Passing a resource ref by value does not fetch
|
||||
content. A helper capability such as `wf.source.read_resource` receives the ref,
|
||||
uses runtime/platform context to resolve the logical source, and applies bounded
|
||||
output policy before returning text into workflow state/output.
|
||||
|
||||
For MCP, the provider also owns stateful upstream sessions:
|
||||
|
||||
```text
|
||||
|
||||
@@ -120,7 +120,11 @@ Provider profiles live in config under `auth.providers`:
|
||||
"token_url": "https://oauth2.googleapis.com/token",
|
||||
"client_id_env": "GOOGLE_OAUTH_CLIENT_ID",
|
||||
"client_secret_env": "GOOGLE_OAUTH_CLIENT_SECRET",
|
||||
"scopes": ["https://www.googleapis.com/auth/drive.readonly"]
|
||||
"scopes": ["https://www.googleapis.com/auth/drive.readonly"],
|
||||
"extra_authorize_params": {
|
||||
"access_type": "offline",
|
||||
"prompt": "consent"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+9
-3
@@ -5,7 +5,7 @@ from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Annotated, Any, Literal, Protocol
|
||||
|
||||
from pydantic import AnyUrl, BaseModel, Field
|
||||
from pydantic import AnyUrl, BaseModel, Field, ValidationError
|
||||
|
||||
AUTH_ID_PATTERN = r"^[A-Za-z0-9_][A-Za-z0-9_.-]*$"
|
||||
|
||||
@@ -149,17 +149,23 @@ def auth_record_from_compat(
|
||||
)
|
||||
if not isinstance(client_id, str) or not client_id:
|
||||
raise ValueError("oauth_refresh_token client_id is required")
|
||||
if not isinstance(client_secret, str):
|
||||
if not isinstance(client_secret, str) or not client_secret:
|
||||
raise ValueError("oauth_refresh_token client_secret is required")
|
||||
if not isinstance(refresh_token, str) or not refresh_token:
|
||||
raise ValueError("oauth_refresh_token refresh_token is required")
|
||||
if not isinstance(token_url, str) or not token_url:
|
||||
raise ValueError("oauth_refresh_token token_url is required")
|
||||
try:
|
||||
validated_token_url = AnyUrl(token_url)
|
||||
except ValidationError as exc:
|
||||
raise ValueError(
|
||||
f"oauth_refresh_token token_url is invalid: {exc}"
|
||||
) from exc
|
||||
auth = OAuthRefreshTokenAuth(
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
refresh_token=refresh_token,
|
||||
token_url=AnyUrl(token_url),
|
||||
token_url=validated_token_url,
|
||||
scopes=scopes,
|
||||
)
|
||||
case _:
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, Protocol
|
||||
|
||||
from wf_platform import page_items
|
||||
|
||||
from .operation_context import WorkflowOperationContext
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class WorkflowSourceDiagnosticsProvider(Protocol):
|
||||
"""Optional source-specific diagnostics provider.
|
||||
@@ -63,7 +66,12 @@ class WorkflowSourceAdminApi:
|
||||
if self.diagnostics is not None:
|
||||
try:
|
||||
payload["diagnostics"] = self.diagnostics.diagnose_source(source_id)
|
||||
except Exception:
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"Source diagnostics failed for source_id=%s: %s",
|
||||
source_id,
|
||||
exc,
|
||||
)
|
||||
payload["diagnostics"] = {
|
||||
"status": "error",
|
||||
"message": "Diagnostics unavailable",
|
||||
|
||||
@@ -41,11 +41,10 @@ async def read_resource(
|
||||
contents = payload.get("contents", [])
|
||||
first = contents[0] if isinstance(contents, list) and contents else {}
|
||||
text = first.get("text") if isinstance(first, dict) else None
|
||||
upstream_truncated = payload.get("truncated") is True
|
||||
truncated = upstream_truncated or (isinstance(text, str) and len(text) > max_chars)
|
||||
if isinstance(text, str) and len(text) > max_chars:
|
||||
text = text[:max_chars]
|
||||
truncated = True
|
||||
else:
|
||||
truncated = False
|
||||
mime_type: str | None = None
|
||||
if isinstance(first, dict) and isinstance(first.get("mimeType"), str):
|
||||
mime_type = first["mimeType"]
|
||||
|
||||
@@ -25,6 +25,9 @@ def validate_deployment_dependencies(
|
||||
for logical_ref, required in artifact.required_capability_map().items():
|
||||
platform_source = sources_by_id.get(required.logical_source)
|
||||
if platform_source is not None and platform_source.platform:
|
||||
# Platform sources have fixed ids matching required.logical_source, so
|
||||
# sources_by_id can resolve them directly and bindings.get(...) is
|
||||
# intentionally bypassed when choosing bound_source_id.
|
||||
bound_source_id = required.logical_source
|
||||
else:
|
||||
bound_source_id = bindings.get(required.logical_source)
|
||||
|
||||
+21
-3
@@ -70,6 +70,16 @@ class OAuthCodeLoginFlow:
|
||||
authorization_response: str | None,
|
||||
authorization_url_callback: Callable[[str, str], str | None] | None = None,
|
||||
) -> OAuthLoginResult:
|
||||
"""Run an OAuth authorization-code login and return durable token data.
|
||||
|
||||
The provider config supplies endpoints, redirect URI, scopes, and any
|
||||
provider-specific authorization parameters. The optional callback lets
|
||||
interactive CLI code display the generated authorization URL and return
|
||||
an out-of-band callback URL; an explicitly supplied authorization
|
||||
response wins over a callback response. `fetch_token` is called only
|
||||
after a response URL is available.
|
||||
"""
|
||||
|
||||
client = self._client_factory(
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
@@ -80,10 +90,11 @@ class OAuthCodeLoginFlow:
|
||||
authorization_url, state = client.create_authorization_url(
|
||||
str(provider.auth_url),
|
||||
redirect_uri=provider.redirect_uri,
|
||||
access_type="offline",
|
||||
prompt="consent",
|
||||
**provider.extra_authorize_params,
|
||||
)
|
||||
if authorization_url_callback is not None:
|
||||
# Interactive and test callbacks can complete the out-of-band flow;
|
||||
# an explicit authorization_response remains the higher-priority input.
|
||||
callback_response = authorization_url_callback(authorization_url, state)
|
||||
if authorization_response is None:
|
||||
authorization_response = callback_response
|
||||
@@ -96,6 +107,13 @@ class OAuthCodeLoginFlow:
|
||||
refresh_token = token.get("refresh_token")
|
||||
if refresh_token is not None and not isinstance(refresh_token, str):
|
||||
raise ValueError("OAuth refresh_token must be a string")
|
||||
subject = token.get("sub")
|
||||
if subject is not None and not isinstance(subject, str):
|
||||
raise ValueError("OAuth sub claim must be a string")
|
||||
raw_scope = token.get("scope")
|
||||
scopes = tuple(str(raw_scope).split()) if raw_scope else provider.scopes
|
||||
return OAuthLoginResult(refresh_token=refresh_token, scopes=scopes)
|
||||
return OAuthLoginResult(
|
||||
refresh_token=refresh_token,
|
||||
subject=subject,
|
||||
scopes=scopes,
|
||||
)
|
||||
|
||||
@@ -216,6 +216,7 @@ class OAuthProviderConfig(WorkflowConfigModel):
|
||||
client_secret_env: str | None = None
|
||||
scopes: tuple[str, ...] = ()
|
||||
redirect_uri: str = "http://127.0.0.1:0/oauth/callback"
|
||||
extra_authorize_params: dict[str, str] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class AuthConfig(WorkflowConfigModel):
|
||||
|
||||
@@ -116,8 +116,34 @@ class ContentAccessService:
|
||||
if resource is None:
|
||||
raise KeyError(f"unknown resource {uri!r} for source {source_id!r}")
|
||||
connection = self.connection_service.get(source_id)
|
||||
return await self.upstream.read_resource(
|
||||
payload = await self.upstream.read_resource(
|
||||
connection,
|
||||
resource.qualified_name,
|
||||
resource.uri,
|
||||
)
|
||||
return _truncate_resource_payload(payload, max_chars=max_chars)
|
||||
|
||||
|
||||
def _truncate_resource_payload(payload: dict[str, Any], *, max_chars: int) -> dict[str, Any]:
|
||||
"""Bound text content returned by source URI reads without mutating upstream payload."""
|
||||
|
||||
contents = payload.get("contents")
|
||||
if not isinstance(contents, list):
|
||||
return payload
|
||||
bounded_contents: list[Any] = []
|
||||
truncated = payload.get("truncated") is True
|
||||
for item in contents:
|
||||
if not isinstance(item, dict) or not isinstance(item.get("text"), str):
|
||||
bounded_contents.append(item)
|
||||
continue
|
||||
text = item["text"]
|
||||
if len(text) <= max_chars:
|
||||
bounded_contents.append(item)
|
||||
continue
|
||||
bounded = dict(item)
|
||||
bounded["text"] = text[:max_chars]
|
||||
bounded_contents.append(bounded)
|
||||
truncated = True
|
||||
if not truncated and bounded_contents == contents:
|
||||
return payload
|
||||
return {**payload, "contents": bounded_contents, "truncated": truncated}
|
||||
|
||||
@@ -180,7 +180,7 @@ class HttpxOAuthTokenRefresher:
|
||||
data["client_secret"] = auth.client_secret
|
||||
if auth.scopes:
|
||||
data["scope"] = " ".join(auth.scopes)
|
||||
async with httpx.AsyncClient() as client:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
response = await client.post(str(auth.token_url), data=data)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
|
||||
@@ -151,3 +151,31 @@ def test_auth_record_from_compat_maps_oauth_refresh_token() -> None:
|
||||
assert record.auth.client_id == "client"
|
||||
assert str(record.auth.token_url) == "https://oauth2.googleapis.com/token"
|
||||
assert record.auth.scopes == ("https://www.googleapis.com/auth/drive.readonly",)
|
||||
|
||||
|
||||
def test_auth_record_from_compat_rejects_bad_oauth_refresh_token_payload() -> None:
|
||||
from wf_api.auth import auth_record_from_compat
|
||||
|
||||
with pytest.raises(ValueError, match="client_secret"):
|
||||
auth_record_from_compat(
|
||||
id="google.drive.personal",
|
||||
scheme="oauth_refresh_token",
|
||||
payload={
|
||||
"client_id": "client",
|
||||
"client_secret": "",
|
||||
"refresh_token": "refresh",
|
||||
"token_url": "https://oauth2.googleapis.com/token",
|
||||
},
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="token_url is invalid"):
|
||||
auth_record_from_compat(
|
||||
id="google.drive.personal",
|
||||
scheme="oauth_refresh_token",
|
||||
payload={
|
||||
"client_id": "client",
|
||||
"client_secret": "secret",
|
||||
"refresh_token": "refresh",
|
||||
"token_url": "not a url",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -39,6 +39,28 @@ async def test_read_resource_resolves_logical_source_and_bounds_text() -> None:
|
||||
assert result.text == "abcde"
|
||||
|
||||
|
||||
async def test_read_resource_preserves_upstream_truncated_signal() -> None:
|
||||
async def handler(source_id: str, uri: str, max_chars: int):
|
||||
return {
|
||||
"contents": [{"type": "text", "text": "abc", "mimeType": "text/plain"}],
|
||||
"truncated": True,
|
||||
}
|
||||
|
||||
platform = SourceBindingPlatformContext(
|
||||
source_bindings={"drive": "drive.personal"},
|
||||
read_resource_handler=handler,
|
||||
)
|
||||
|
||||
result = await read_resource(
|
||||
SourceResourceRef(logical_source="drive", uri="gdrive://file/abc"),
|
||||
RuntimeContext(current_node_id="read", platform=platform),
|
||||
max_chars=5,
|
||||
)
|
||||
|
||||
assert result.truncated is True
|
||||
assert result.text == "abc"
|
||||
|
||||
|
||||
async def test_read_resource_requires_platform_context() -> None:
|
||||
with pytest.raises(RuntimeError, match="platform context"):
|
||||
await read_resource(
|
||||
|
||||
@@ -15,6 +15,7 @@ from wf_config import OAuthProviderConfig
|
||||
def _oauth_provider(
|
||||
*,
|
||||
scopes: tuple[str, ...] = (),
|
||||
extra_authorize_params: dict[str, str] | None = None,
|
||||
) -> OAuthProviderConfig:
|
||||
return OAuthProviderConfig(
|
||||
kind="oauth_authorization_code_pkce",
|
||||
@@ -23,6 +24,7 @@ def _oauth_provider(
|
||||
client_id_env="GOOGLE_OAUTH_CLIENT_ID",
|
||||
client_secret_env="GOOGLE_OAUTH_CLIENT_SECRET",
|
||||
scopes=scopes,
|
||||
extra_authorize_params=extra_authorize_params or {},
|
||||
)
|
||||
|
||||
|
||||
@@ -85,12 +87,14 @@ class _FakeOAuthClient:
|
||||
return {
|
||||
"refresh_token": "refresh",
|
||||
"scope": "https://www.googleapis.com/auth/drive.readonly",
|
||||
"sub": "user-123",
|
||||
}
|
||||
|
||||
|
||||
async def test_oauth_code_login_flow_uses_injected_client() -> None:
|
||||
provider = _oauth_provider(
|
||||
scopes=("https://www.googleapis.com/auth/drive.readonly",),
|
||||
extra_authorize_params={"access_type": "offline", "prompt": "consent"},
|
||||
)
|
||||
clients: list[_FakeOAuthClient] = []
|
||||
|
||||
@@ -109,6 +113,7 @@ async def test_oauth_code_login_flow_uses_injected_client() -> None:
|
||||
)
|
||||
|
||||
assert result.refresh_token == "refresh"
|
||||
assert result.subject == "user-123"
|
||||
assert result.scopes == ("https://www.googleapis.com/auth/drive.readonly",)
|
||||
client = clients[0]
|
||||
assert client.init_kwargs["redirect_uri"] == provider.redirect_uri
|
||||
@@ -174,6 +179,10 @@ def test_auth_oauth_login_saves_record_from_provider_profile(monkeypatch, tmp_pa
|
||||
"token_url": "https://oauth2.googleapis.com/token",
|
||||
"client_id_env": "GOOGLE_OAUTH_CLIENT_ID",
|
||||
"scopes": ["https://www.googleapis.com/auth/drive.readonly"],
|
||||
"extra_authorize_params": {
|
||||
"access_type": "offline",
|
||||
"prompt": "consent",
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -472,6 +472,10 @@ def test_workflow_config_parses_oauth_provider_profile() -> None:
|
||||
"scopes": [
|
||||
"https://www.googleapis.com/auth/drive.readonly",
|
||||
],
|
||||
"extra_authorize_params": {
|
||||
"access_type": "offline",
|
||||
"prompt": "consent",
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -482,3 +486,7 @@ def test_workflow_config_parses_oauth_provider_profile() -> None:
|
||||
assert provider.kind == "oauth_authorization_code_pkce"
|
||||
assert provider.client_id_env == "GOOGLE_OAUTH_CLIENT_ID"
|
||||
assert provider.scopes == ("https://www.googleapis.com/auth/drive.readonly",)
|
||||
assert provider.extra_authorize_params == {
|
||||
"access_type": "offline",
|
||||
"prompt": "consent",
|
||||
}
|
||||
|
||||
@@ -286,6 +286,26 @@ async def test_read_resource_by_source_uri_reads_upstream() -> None:
|
||||
assert result["contents"][0]["text"] == "Welcome from the fake adapter resource."
|
||||
|
||||
|
||||
async def test_read_resource_by_source_uri_bounds_text() -> None:
|
||||
service = WfMcpService(
|
||||
store=FileStore(local_temp_root() / "content_source_uri_bound")
|
||||
)
|
||||
service.register_connection(
|
||||
ConnectionConfig(id="demo.personal", server="demo", account="personal")
|
||||
)
|
||||
service.register_adapter("demo", FakeAdapter())
|
||||
await service.refresh_connection_catalog("demo.personal")
|
||||
|
||||
result = await service.content_access.read_resource_by_source_uri(
|
||||
source_id="demo.personal",
|
||||
uri="demo://docs/welcome",
|
||||
max_chars=7,
|
||||
)
|
||||
|
||||
assert result["contents"][0]["text"] == "Welcome"
|
||||
assert result["truncated"] is True
|
||||
|
||||
|
||||
async def test_read_resource_by_source_uri_rejects_unknown_resource() -> None:
|
||||
service = WfMcpService(
|
||||
store=FileStore(local_temp_root() / "content_source_uri_unknown")
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from wf_mcp.broker.service.source_diagnostics import SourceDiagnosticsProvider
|
||||
from wf_mcp.connections import ConnectionRegistry
|
||||
from wf_mcp.models import ConnectionConfig
|
||||
@@ -18,7 +20,7 @@ def _connection(**metadata: object) -> ConnectionConfig:
|
||||
)
|
||||
|
||||
|
||||
def _provider(tmp_path, connection: ConnectionConfig) -> SourceDiagnosticsProvider:
|
||||
def _provider(tmp_path: Path, connection: ConnectionConfig) -> SourceDiagnosticsProvider:
|
||||
registry = ConnectionRegistry()
|
||||
registry.register(connection)
|
||||
return SourceDiagnosticsProvider(
|
||||
|
||||
@@ -97,6 +97,9 @@ async def test_httpx_oauth_refresher_posts_refresh_token_grant(
|
||||
return {"access_token": "access-token", "expires_in": 3600}
|
||||
|
||||
class _Client:
|
||||
def __init__(self, **kwargs: object) -> None:
|
||||
assert kwargs["timeout"] == 10.0
|
||||
|
||||
async def __aenter__(self) -> "_Client":
|
||||
return self
|
||||
|
||||
|
||||
Reference in New Issue
Block a user