From 7f3f219f58b86bde94ee8f7dc1c76ab0da4e009f Mon Sep 17 00:00:00 2001 From: wizardchen Date: Mon, 22 Jun 2026 20:15:52 +0800 Subject: [PATCH] feat(mcp): add OAuth2 authorization for MCP services Support connecting to MCP servers that require OAuth2, with a Cursor-like zero-config flow: the user provides only a URL and the system performs metadata discovery, dynamic client registration (DCR), and the authorization code flow with PKCE on first connect. - Persist per-user access/refresh tokens and per-service DCR client credentials, encrypted at rest (AES-256-GCM) - Store the PKCE verifier/state during the flow via Redis with an in-memory fallback for Lite mode - Cache MCP clients per (service, user) so tokens stay user-isolated - Add endpoints for authorize-url, callback, status, and token revoke; the callback is allowlisted for unauthenticated access - Surface a structured "authorization required" hint when a tool call hits an unauthorized MCP service - Frontend: add an auth-type selector with OAuth config, inline authorization status/actions, and render the connection test result inline in the editor drawer instead of a stacked centered dialog --- frontend/src/api/mcp-service.ts | 40 ++ frontend/src/i18n/locales/zh-CN.ts | 18 + frontend/src/views/settings/McpSettings.vue | 26 +- .../settings/components/McpServiceDialog.vue | 317 ++++++++-- .../settings/components/McpTestResult.vue | 563 ------------------ .../settings/components/McpTestResultBody.vue | 373 ++++++++++++ internal/agent/tools/mcp_tool.go | 49 +- internal/application/repository/mcp_oauth.go | 90 +++ internal/application/service/mcp_service.go | 46 +- .../application/service/mcp_service_test.go | 3 +- internal/container/container.go | 3 + internal/handler/dto/mcp.go | 14 +- internal/handler/mcp_oauth.go | 204 +++++++ internal/handler/mcp_service.go | 17 + internal/mcp/client.go | 81 ++- internal/mcp/manager.go | 79 ++- internal/mcp/oauth_manager.go | 189 ++++++ internal/mcp/oauth_state.go | 125 ++++ internal/mcp/oauth_tokenstore.go | 70 +++ internal/middleware/auth.go | 4 + internal/router/router.go | 16 +- internal/types/interfaces/mcp_oauth.go | 32 + internal/types/mcp.go | 35 ++ internal/types/mcp_oauth.go | 139 +++++ migrations/sqlite/000000_init.up.sql | 33 + .../versioned/000062_mcp_oauth.down.sql | 2 + migrations/versioned/000062_mcp_oauth.up.sql | 38 ++ 27 files changed, 1920 insertions(+), 686 deletions(-) delete mode 100644 frontend/src/views/settings/components/McpTestResult.vue create mode 100644 frontend/src/views/settings/components/McpTestResultBody.vue create mode 100644 internal/application/repository/mcp_oauth.go create mode 100644 internal/handler/mcp_oauth.go create mode 100644 internal/mcp/oauth_manager.go create mode 100644 internal/mcp/oauth_state.go create mode 100644 internal/mcp/oauth_tokenstore.go create mode 100644 internal/types/interfaces/mcp_oauth.go create mode 100644 internal/types/mcp_oauth.go create mode 100644 migrations/versioned/000062_mcp_oauth.down.sql create mode 100644 migrations/versioned/000062_mcp_oauth.up.sql diff --git a/frontend/src/api/mcp-service.ts b/frontend/src/api/mcp-service.ts index 8af1a1a74..4404690f3 100644 --- a/frontend/src/api/mcp-service.ts +++ b/frontend/src/api/mcp-service.ts @@ -10,6 +10,10 @@ export interface MCPService { url?: string // Optional: required for SSE/HTTP Streamable headers?: Record auth_config?: { + // Authentication strategy. Empty/absent means none. "oauth" enables the + // per-user OAuth2 authorization-code flow (zero-config: discovery + + // dynamic client registration). + auth_type?: '' | 'api_key' | 'bearer' | 'oauth' // Secret fields (api_key, token) are NEVER returned by the server in // this shape — they live behind the /credentials subresource. The // optional-property typing remains so create-mode payloads can still @@ -17,6 +21,9 @@ export interface MCPService { api_key?: string token?: string custom_headers?: Record + // OAuth-only, non-secret configuration. + scopes?: string[] + auth_server_metadata_url?: string } advanced_config?: { timeout?: number @@ -166,6 +173,39 @@ export async function deleteMCPCredentialField( await del(`/api/v1/mcp-services/${serviceId}/credentials/${field}`) } +// ---------------------------------------------------------------------------- +// Per-user OAuth2 authorization-code flow. +// +// The user authorizes a service once; the backend stores their access/refresh +// token (per tenant + user + service) and refreshes it transparently. The +// callback is a public backend route that the third-party authorization +// server redirects to. +// ---------------------------------------------------------------------------- + +// Path of the public backend OAuth callback (registered outside /mcp-services +// to avoid a route conflict, and allow-listed for no-auth in the backend). +export const MCP_OAUTH_CALLBACK_PATH = '/api/v1/mcp-oauth/callback' + +// Begin authorization for the current user. Returns the URL to open in a popup. +export async function getMCPOAuthAuthorizeURL( + serviceId: string, + body: { redirect_uri: string; frontend_redirect?: string } +): Promise { + const response: any = await post(`/api/v1/mcp-services/${serviceId}/oauth/authorize-url`, body) + return (response.data ?? response)?.authorization_url ?? '' +} + +// Whether the current user has authorized this service. +export async function getMCPOAuthStatus(serviceId: string): Promise { + const response: any = await get(`/api/v1/mcp-services/${serviceId}/oauth/status`) + return Boolean((response.data ?? response)?.authorized) +} + +// Revoke the current user's token (forces re-authorization). +export async function revokeMCPOAuthToken(serviceId: string): Promise { + await del(`/api/v1/mcp-services/${serviceId}/oauth/token`) +} + export async function resolveToolApproval( pendingId: string, body: { decision: 'approve' | 'reject'; modified_args?: Record; reason?: string } diff --git a/frontend/src/i18n/locales/zh-CN.ts b/frontend/src/i18n/locales/zh-CN.ts index 67ddd2d15..794a1b624 100755 --- a/frontend/src/i18n/locales/zh-CN.ts +++ b/frontend/src/i18n/locales/zh-CN.ts @@ -4343,6 +4343,7 @@ export default { connectionSection: "连接配置", enableServiceDesc: "关闭后该服务不会被调用", testAfterSaveHint: "保存后可测试连接", + testResultTitle: "测试结果", unitSecond: "秒", unitTimes: "次", name: "服务名称", @@ -4367,6 +4368,19 @@ export default { addEnvVar: "添加环境变量", enableService: "启用服务", authConfig: "认证配置", + authType: "认证方式", + authTypeNone: "无 / 自定义 Header", + authTypeApiKey: "API Key", + authTypeBearer: "Bearer Token", + authTypeOAuth: "OAuth 2.0(首次连接授权)", + oauthScopes: "Scopes(可选,空格分隔)", + oauthAuthorization: "授权状态", + oauthAuthorized: "已授权", + oauthUnauthorized: "未授权", + oauthAuthorize: "去授权", + oauthReauthorize: "重新授权", + oauthRevoke: "撤销授权", + oauthSaveFirstHint: "保存服务后,可在编辑页发起首次授权(每个用户独立授权)。", apiKey: "API Key", bearerToken: "Bearer Token", optional: "可选", @@ -4387,6 +4401,10 @@ export default { updated: "MCP 服务已更新", createFailed: "创建 MCP 服务失败", updateFailed: "更新 MCP 服务失败", + authorized: "授权成功", + authorizeFailed: "发起授权失败", + revoked: "已撤销授权", + revokeFailed: "撤销失败", }, }, promptTemplate: { diff --git a/frontend/src/views/settings/McpSettings.vue b/frontend/src/views/settings/McpSettings.vue index 2cdca4b1a..70a356c37 100644 --- a/frontend/src/views/settings/McpSettings.vue +++ b/frontend/src/views/settings/McpSettings.vue @@ -112,15 +112,6 @@ :service="currentService" :mode="dialogMode" @success="handleDialogSuccess" - @test="handleDrawerTest" - /> - - - @@ -134,11 +125,9 @@ import { listMCPServices, updateMCPService, deleteMCPService, - type MCPService, - type MCPTestResult + type MCPService } from '@/api/mcp-service' import McpServiceDialog from './components/McpServiceDialog.vue' -import McpTestResult from './components/McpTestResult.vue' import { useConfirmDelete } from '@/components/settings/useConfirmDelete' import { useAuthStore } from '@/stores/auth' @@ -151,10 +140,6 @@ const loading = ref(false) const dialogVisible = ref(false) const dialogMode = ref<'add' | 'edit'>('add') const currentService = ref(null) -const testDialogVisible = ref(false) -const testResult = ref(null) -const testingServiceName = ref('') -const testingServiceId = ref('') // Load MCP services const loadServices = async () => { @@ -266,15 +251,6 @@ const getBuiltinServiceOptions = () => { ] } -// Drawer 内点击"测试连接"后,复用现有的 testResult dialog 展示结果。 -// 抽屉只负责调 API + 拿结果,弹窗位置/状态由父组件统一管。 -const handleDrawerTest = ({ service, result }: { service: MCPService; result: MCPTestResult }) => { - testingServiceName.value = service.name - testingServiceId.value = service.id - testResult.value = result - testDialogVisible.value = true -} - // Handle menu action. 'test' has been removed from the menu — testing now // lives only in the editor drawer. We keep the switch's case list narrow // so a stray 'test' from somewhere else falls through harmlessly. diff --git a/frontend/src/views/settings/components/McpServiceDialog.vue b/frontend/src/views/settings/components/McpServiceDialog.vue index 4fa20d2f8..f2cbe1e0f 100644 --- a/frontend/src/views/settings/components/McpServiceDialog.vue +++ b/frontend/src/views/settings/components/McpServiceDialog.vue @@ -123,43 +123,88 @@ - +

{{ t('mcpServiceDialog.authConfig') }}

+
+ + +
+ + + + -
@@ -217,12 +262,30 @@ + + +
+
+

{{ t('mcpServiceDialog.testResultTitle', '测试结果') }}

+ + + +
+ +
- - - diff --git a/frontend/src/views/settings/components/McpTestResultBody.vue b/frontend/src/views/settings/components/McpTestResultBody.vue new file mode 100644 index 000000000..db0961b82 --- /dev/null +++ b/frontend/src/views/settings/components/McpTestResultBody.vue @@ -0,0 +1,373 @@ + + + + + diff --git a/internal/agent/tools/mcp_tool.go b/internal/agent/tools/mcp_tool.go index 07b7273bc..a4eaa1f8f 100644 --- a/internal/agent/tools/mcp_tool.go +++ b/internal/agent/tools/mcp_tool.go @@ -11,6 +11,7 @@ import ( "github.com/Tencent/WeKnora/internal/logger" "github.com/Tencent/WeKnora/internal/mcp" "github.com/Tencent/WeKnora/internal/types" + mcpclient "github.com/mark3labs/mcp-go/client" ) type MCPInput = map[string]any @@ -173,12 +174,12 @@ func (t *MCPTool) Execute(ctx context.Context, args json.RawMessage) (*types.Too } // Get or create MCP client - client, err := t.mcpManager.GetOrCreateClient(t.service) + client, err := t.mcpManager.GetOrCreateClient(ctx, t.service) if err != nil { logger.GetLogger(ctx).Errorf("Failed to get MCP client: %v", err) return &types.ToolResult{ Success: false, - Error: fmt.Sprintf("Failed to connect to MCP service: %v", err), + Error: oauthAwareConnectError(t.service, err), }, nil } @@ -200,12 +201,12 @@ func (t *MCPTool) Execute(ctx context.Context, args json.RawMessage) (*types.Too logger.GetLogger(ctx).Warnf("MCP tool call failed, retrying with fresh connection: %v", err) _ = client.Disconnect() - client, err = t.mcpManager.GetOrCreateClient(t.service) + client, err = t.mcpManager.GetOrCreateClient(ctx, t.service) if err != nil { logger.GetLogger(ctx).Errorf("Failed to reconnect MCP client: %v", err) return &types.ToolResult{ Success: false, - Error: fmt.Sprintf("Failed to reconnect to MCP service: %v", err), + Error: oauthAwareConnectError(t.service, err), }, nil } @@ -215,7 +216,7 @@ func (t *MCPTool) Execute(ctx context.Context, args json.RawMessage) (*types.Too logger.GetLogger(ctx).Errorf("MCP tool call failed: %v", err) return &types.ToolResult{ Success: false, - Error: fmt.Sprintf("Tool execution failed: %v", err), + Error: oauthAwareConnectError(t.service, err), }, nil } @@ -374,6 +375,38 @@ func extractContentText(content []mcp.ContentItem) string { return strings.Join(textParts, "\n") } +// oauthAwareConnectError turns a low-level MCP connect/call error into a +// message the agent (and ultimately the user) can act on. For OAuth services +// that have not been authorized — or whose token expired and could not be +// refreshed — the underlying library surfaces an authorization-required error; +// we translate that into an explicit instruction to authorize, instead of an +// opaque "connection failed". +func oauthAwareConnectError(service *types.MCPService, err error) string { + if service.AuthConfig.IsOAuth() && isAuthorizationRequired(err) { + return fmt.Sprintf( + "MCP service %q requires OAuth authorization. Please open the service settings "+ + "and click \"Authorize\" to grant access, then retry.", + service.Name, + ) + } + return fmt.Sprintf("Failed to connect to MCP service: %v", err) +} + +// isAuthorizationRequired reports whether err indicates the OAuth flow has not +// been completed (no valid token / 401). +func isAuthorizationRequired(err error) bool { + if err == nil { + return false + } + if mcpclient.IsOAuthAuthorizationRequiredError(err) || mcpclient.IsAuthorizationRequiredError(err) { + return true + } + msg := err.Error() + return strings.Contains(msg, "authorization required") || + strings.Contains(msg, "no valid token") || + strings.Contains(msg, "401") +} + // sanitizeName sanitizes a name to create a valid identifier func sanitizeName(name string) string { // Replace invalid characters with underscores @@ -422,7 +455,7 @@ func RegisterMCPTools( } // Get or create client (this may take time, but has its own timeout) - client, err := mcpManager.GetOrCreateClient(service) + client, err := mcpManager.GetOrCreateClient(ctx, service) if err != nil { logger.GetLogger(ctx).Errorf("Failed to create MCP client for service %s: %v", service.Name, err) continue @@ -448,7 +481,7 @@ func RegisterMCPTools( logger.GetLogger(ctx).Warnf("Failed to list tools from MCP service %s (will retry with fresh connection): %v", service.Name, err) _ = client.Disconnect() - client, err = mcpManager.GetOrCreateClient(service) + client, err = mcpManager.GetOrCreateClient(ctx, service) if err != nil { logger.GetLogger(ctx).Errorf("Failed to reconnect MCP client for service %s: %v", service.Name, err) continue @@ -504,7 +537,7 @@ func GetMCPToolsInfo( continue } - client, err := mcpManager.GetOrCreateClient(service) + client, err := mcpManager.GetOrCreateClient(ctx, service) if err != nil { continue } diff --git a/internal/application/repository/mcp_oauth.go b/internal/application/repository/mcp_oauth.go new file mode 100644 index 000000000..afd3247ef --- /dev/null +++ b/internal/application/repository/mcp_oauth.go @@ -0,0 +1,90 @@ +package repository + +import ( + "context" + "errors" + "time" + + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +// mcpOAuthRepository implements interfaces.MCPOAuthRepository. +type mcpOAuthRepository struct { + db *gorm.DB +} + +// NewMCPOAuthRepository creates a new MCP OAuth repository. +func NewMCPOAuthRepository(db *gorm.DB) interfaces.MCPOAuthRepository { + return &mcpOAuthRepository{db: db} +} + +func (r *mcpOAuthRepository) GetClient( + ctx context.Context, tenantID uint64, serviceID string, +) (*types.MCPOAuthClient, error) { + var client types.MCPOAuthClient + err := r.db.WithContext(ctx). + Where("tenant_id = ? AND service_id = ?", tenantID, serviceID). + First(&client).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, err + } + return &client, nil +} + +func (r *mcpOAuthRepository) SaveClient(ctx context.Context, client *types.MCPOAuthClient) error { + client.UpdatedAt = time.Now() + return r.db.WithContext(ctx). + Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "tenant_id"}, {Name: "service_id"}}, + DoUpdates: clause.AssignmentColumns([]string{"client_id", "client_secret", "redirect_uri", "updated_at"}), + }). + Create(client).Error +} + +func (r *mcpOAuthRepository) DeleteClient(ctx context.Context, tenantID uint64, serviceID string) error { + return r.db.WithContext(ctx). + Where("tenant_id = ? AND service_id = ?", tenantID, serviceID). + Delete(&types.MCPOAuthClient{}).Error +} + +func (r *mcpOAuthRepository) GetToken( + ctx context.Context, tenantID uint64, userID, serviceID string, +) (*types.MCPOAuthToken, error) { + var token types.MCPOAuthToken + err := r.db.WithContext(ctx). + Where("tenant_id = ? AND user_id = ? AND service_id = ?", tenantID, userID, serviceID). + First(&token).Error + if err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + return nil, err + } + return &token, nil +} + +func (r *mcpOAuthRepository) SaveToken(ctx context.Context, token *types.MCPOAuthToken) error { + token.UpdatedAt = time.Now() + return r.db.WithContext(ctx). + Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "tenant_id"}, {Name: "user_id"}, {Name: "service_id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "access_token", "refresh_token", "token_type", "expires_at", "updated_at", + }), + }). + Create(token).Error +} + +func (r *mcpOAuthRepository) DeleteToken( + ctx context.Context, tenantID uint64, userID, serviceID string, +) error { + return r.db.WithContext(ctx). + Where("tenant_id = ? AND user_id = ? AND service_id = ?", tenantID, userID, serviceID). + Delete(&types.MCPOAuthToken{}).Error +} diff --git a/internal/application/service/mcp_service.go b/internal/application/service/mcp_service.go index 45dca4d8d..e40848dad 100644 --- a/internal/application/service/mcp_service.go +++ b/internal/application/service/mcp_service.go @@ -18,16 +18,19 @@ import ( type mcpServiceService struct { mcpServiceRepo interfaces.MCPServiceRepository mcpManager *mcp.MCPManager + oauthRepo interfaces.MCPOAuthRepository } // NewMCPServiceService creates a new MCP service service func NewMCPServiceService( mcpServiceRepo interfaces.MCPServiceRepository, mcpManager *mcp.MCPManager, + oauthRepo interfaces.MCPOAuthRepository, ) interfaces.MCPServiceService { return &mcpServiceService{ mcpServiceRepo: mcpServiceRepo, mcpManager: mcpManager, + oauthRepo: oauthRepo, } } @@ -168,14 +171,36 @@ func (s *mcpServiceService) UpdateMCPService(ctx context.Context, service *types maps.Copy(preHeaders, existing.AuthConfig.CustomHeaders) } + preAuthType := types.MCPAuthNone + if existing.AuthConfig != nil { + preAuthType = existing.AuthConfig.AuthType + } + // CustomHeaders flows through main PUT (it's structural, not a secret) — // nil preserves, non-nil replaces. Other AuthConfig fields (APIKey/Token) // are never accepted via main PUT; the handler strips them up front. - if service.AuthConfig != nil && service.AuthConfig.CustomHeaders != nil { + // + // auth_type / scopes / auth_server_metadata_url are non-secret OAuth + // configuration and also flow through here. + if service.AuthConfig != nil { if existing.AuthConfig == nil { existing.AuthConfig = &types.MCPAuthConfig{} } - existing.AuthConfig.CustomHeaders = service.AuthConfig.CustomHeaders + if service.AuthConfig.CustomHeaders != nil { + existing.AuthConfig.CustomHeaders = service.AuthConfig.CustomHeaders + } + // Only overwrite OAuth config when explicitly provided, so a partial + // PUT that carries only custom_headers does not wipe an existing + // auth_type / scopes. (Empty/absent is treated as "no change".) + if service.AuthConfig.AuthType != types.MCPAuthNone { + existing.AuthConfig.AuthType = service.AuthConfig.AuthType + } + if service.AuthConfig.Scopes != nil { + existing.AuthConfig.Scopes = service.AuthConfig.Scopes + } + if service.AuthConfig.AuthServerMetadataURL != "" { + existing.AuthConfig.AuthServerMetadataURL = service.AuthConfig.AuthServerMetadataURL + } } // Merge updates: only update fields that are provided (non-zero or explicitly set) @@ -257,6 +282,9 @@ func (s *mcpServiceService) UpdateMCPService(ctx context.Context, service *types if !maps.Equal(currHeaders, preHeaders) { configChanged = true } + if existing.AuthConfig != nil && existing.AuthConfig.AuthType != preAuthType { + configChanged = true + } name := secutils.SanitizeForLog(existing.Name) // Close existing client connection if: // 1. Service is now disabled (need to close connection) @@ -321,10 +349,18 @@ func (s *mcpServiceService) TestMCPService( return nil, fmt.Errorf("MCP service not found") } - // Create temporary client for testing + // Create temporary client for testing. For OAuth services, wire the + // per-user token store so the test connects with the current user's + // authorization (and surfaces an authorization-required message when the + // user has not authorized yet). config := &mcp.ClientConfig{ Service: service, } + if service.AuthConfig.IsOAuth() { + config.OAuthRepo = s.oauthRepo + config.TenantID, _ = types.TenantIDFromContext(ctx) + config.UserID, _ = types.UserIDFromContext(ctx) + } client, err := mcp.NewMCPClient(config) if err != nil { @@ -397,7 +433,7 @@ func (s *mcpServiceService) GetMCPServiceTools( } // Get or create client - client, err := s.mcpManager.GetOrCreateClient(service) + client, err := s.mcpManager.GetOrCreateClient(ctx, service) if err != nil { return nil, fmt.Errorf("failed to get MCP client: %w", err) } @@ -535,7 +571,7 @@ func (s *mcpServiceService) GetMCPServiceResources( } // Get or create client - client, err := s.mcpManager.GetOrCreateClient(service) + client, err := s.mcpManager.GetOrCreateClient(ctx, service) if err != nil { return nil, fmt.Errorf("failed to get MCP client: %w", err) } diff --git a/internal/application/service/mcp_service_test.go b/internal/application/service/mcp_service_test.go index a040d11cf..56750e626 100644 --- a/internal/application/service/mcp_service_test.go +++ b/internal/application/service/mcp_service_test.go @@ -111,7 +111,8 @@ func newTestService() (*mcpServiceService, *fakeMCPRepo) { repo := newFakeMCPRepo() svc := &mcpServiceService{ mcpServiceRepo: repo, - mcpManager: mcp.NewMCPManager(), + mcpManager: mcp.NewMCPManager(nil), + oauthRepo: nil, } return svc, repo } diff --git a/internal/container/container.go b/internal/container/container.go index 22727ee48..f3c772684 100644 --- a/internal/container/container.go +++ b/internal/container/container.go @@ -153,6 +153,7 @@ func BuildContainer(container *dig.Container) *dig.Container { must(container.Provide(memoryRepo.NewMemoryRepository)) must(container.Provide(repository.NewMCPServiceRepository)) must(container.Provide(repository.NewMCPToolApprovalRepository)) + must(container.Provide(repository.NewMCPOAuthRepository)) must(container.Provide(repository.NewCustomAgentRepository)) must(container.Provide(repository.NewOrganizationRepository)) must(container.Provide(repository.NewKBShareRepository)) @@ -171,6 +172,7 @@ func BuildContainer(container *dig.Container) *dig.Container { // MCP manager for managing MCP client connections logger.Debugf(ctx, "[Container] Registering MCP manager...") must(container.Provide(mcp.NewMCPManager)) + must(container.Provide(mcp.NewOAuthManager)) // Business service layer logger.Debugf(ctx, "[Container] Registering business services...") @@ -325,6 +327,7 @@ func BuildContainer(container *dig.Container) *dig.Container { must(container.Provide(handler.NewSystemHandler)) must(container.Provide(handler.NewMCPServiceHandler)) must(container.Provide(handler.NewMCPCredentialsHandler)) + must(container.Provide(handler.NewMCPOAuthHandler)) must(container.Provide(handler.NewModelCredentialsHandler)) must(container.Provide(handler.NewWebSearchProviderCredentialsHandler)) must(container.Provide(handler.NewDataSourceCredentialsHandler)) diff --git a/internal/handler/dto/mcp.go b/internal/handler/dto/mcp.go index 5ba2ae89b..dba340e35 100644 --- a/internal/handler/dto/mcp.go +++ b/internal/handler/dto/mcp.go @@ -45,9 +45,14 @@ type MCPServiceResponse struct { } // MCPAuthConfigResponse intentionally has no APIKey or Token fields. Their -// presence is signalled via MCPServiceResponse.Credentials. +// presence is signalled via MCPServiceResponse.Credentials. AuthType, Scopes +// and AuthServerMetadataURL are non-secret OAuth configuration and are safe to +// echo back so the UI can render the current strategy. type MCPAuthConfigResponse struct { - CustomHeaders map[string]string `json:"custom_headers,omitempty"` + AuthType types.MCPAuthType `json:"auth_type,omitempty"` + CustomHeaders map[string]string `json:"custom_headers,omitempty"` + Scopes []string `json:"scopes,omitempty"` + AuthServerMetadataURL string `json:"auth_server_metadata_url,omitempty"` } // CredentialFieldMetadata reports whether a credential field has a value @@ -84,7 +89,10 @@ func NewMCPServiceResponse(svc *types.MCPService) *MCPServiceResponse { } if svc.AuthConfig != nil { resp.AuthConfig = &MCPAuthConfigResponse{ - CustomHeaders: svc.AuthConfig.CustomHeaders, + AuthType: svc.AuthConfig.AuthType, + CustomHeaders: svc.AuthConfig.CustomHeaders, + Scopes: svc.AuthConfig.Scopes, + AuthServerMetadataURL: svc.AuthConfig.AuthServerMetadataURL, } } if svc.IsBuiltin { diff --git a/internal/handler/mcp_oauth.go b/internal/handler/mcp_oauth.go new file mode 100644 index 000000000..dd1057d88 --- /dev/null +++ b/internal/handler/mcp_oauth.go @@ -0,0 +1,204 @@ +package handler + +import ( + "net/http" + "strings" + + "github.com/Tencent/WeKnora/internal/errors" + "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/mcp" + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" + secutils "github.com/Tencent/WeKnora/internal/utils" + "github.com/gin-gonic/gin" +) + +// MCPOAuthHandler exposes the per-user MCP OAuth2 authorization-code flow: +// kicking off authorization (discovery + dynamic client registration + PKCE), +// receiving the provider redirect, reporting authorization status, and +// revoking a user's token. +type MCPOAuthHandler struct { + oauth *mcp.OAuthManager + mcpManager *mcp.MCPManager + svc interfaces.MCPServiceService +} + +// NewMCPOAuthHandler constructs the handler. +func NewMCPOAuthHandler( + oauth *mcp.OAuthManager, + mcpManager *mcp.MCPManager, + svc interfaces.MCPServiceService, +) *MCPOAuthHandler { + return &MCPOAuthHandler{oauth: oauth, mcpManager: mcpManager, svc: svc} +} + +type mcpOAuthAuthorizeRequest struct { + // RedirectURI is the absolute backend callback URL registered with the + // authorization server (e.g. https://host/api/v1/mcp-services/oauth/callback). + RedirectURI string `json:"redirect_uri"` + // FrontendRedirect is where the callback bounces the browser when done + // (e.g. the MCP settings page). Optional; defaults to "/". + FrontendRedirect string `json:"frontend_redirect"` +} + +// AuthorizeURL begins authorization and returns the URL the browser must open. +// +// AuthorizeURL godoc +// @Summary 发起 MCP OAuth 授权 +// @Description 对使用 OAuth 的 MCP 服务执行发现与动态客户端注册,返回浏览器应跳转的授权地址(当前用户维度) +// @Tags MCP服务 +// @Accept json +// @Produce json +// @Param id path string true "MCP 服务 ID" +// @Param request body map[string]interface{} true "{redirect_uri: string, frontend_redirect?: string}" +// @Success 200 {object} map[string]interface{} "{authorization_url: string}" +// @Failure 400 {object} errors.AppError +// @Security Bearer +// @Router /mcp-services/{id}/oauth/authorize-url [post] +func (h *MCPOAuthHandler) AuthorizeURL(c *gin.Context) { + ctx := c.Request.Context() + serviceID := c.Param("id") + tenantID := c.GetUint64(types.TenantIDContextKey.String()) + userID, _ := types.UserIDFromContext(ctx) + if tenantID == 0 || userID == "" { + c.Error(errors.NewUnauthorizedError("authentication required")) + return + } + + var req mcpOAuthAuthorizeRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.Error(errors.NewBadRequestError(err.Error())) + return + } + req.RedirectURI = strings.TrimSpace(req.RedirectURI) + if req.RedirectURI == "" { + c.Error(errors.NewValidationError("redirect_uri is required")) + return + } + if req.FrontendRedirect == "" { + req.FrontendRedirect = "/" + } + + service, err := h.svc.GetMCPServiceByID(ctx, tenantID, serviceID) + if err != nil || service == nil { + c.Error(errors.NewNotFoundError("MCP service not found")) + return + } + if !service.AuthConfig.IsOAuth() { + c.Error(errors.NewValidationError("MCP service is not configured to use OAuth")) + return + } + + authURL, err := h.oauth.StartAuthorization(ctx, service, tenantID, userID, req.RedirectURI, req.FrontendRedirect) + if err != nil { + logger.ErrorWithFields(ctx, err, map[string]interface{}{ + "service_id": secutils.SanitizeForLog(serviceID), + }) + c.Error(errors.NewInternalServerError("failed to start authorization: " + err.Error())) + return + } + + c.JSON(http.StatusOK, gin.H{"success": true, "data": gin.H{"authorization_url": authURL}}) +} + +// Callback receives the authorization-server redirect. It is registered as a +// public (no-bearer) route; the opaque single-use `state` parameter +// authenticates the request. On completion it redirects the browser back to +// the frontend with the result encoded in the URL fragment. +// +// Callback godoc +// @Summary MCP OAuth 回调 +// @Description 接收授权服务器回调并完成 code 交换,随后重定向回前端 +// @Tags MCP服务 +// @Param code query string false "授权码" +// @Param state query string false "状态参数" +// @Param error query string false "授权错误码" +// @Success 302 +// @Router /mcp-services/oauth/callback [get] +func (h *MCPOAuthHandler) Callback(c *gin.Context) { + ctx := c.Request.Context() + state := strings.TrimSpace(c.Query("state")) + code := strings.TrimSpace(c.Query("code")) + providerErr := strings.TrimSpace(c.Query("error")) + + const fallbackRedirect = "/" + + if providerErr != "" { + c.Redirect(http.StatusFound, fallbackRedirect+"#mcp_oauth_error="+urlQueryEscape(providerErr)) + return + } + if state == "" || code == "" { + c.Redirect(http.StatusFound, fallbackRedirect+"#mcp_oauth_error="+urlQueryEscape("missing_code_or_state")) + return + } + + frontendRedirect, err := h.oauth.CompleteAuthorization(ctx, state, code) + if frontendRedirect == "" { + frontendRedirect = fallbackRedirect + } + if err != nil { + logger.Errorf(ctx, "MCP OAuth callback failed: %v", err) + c.Redirect(http.StatusFound, frontendRedirect+"#mcp_oauth_error="+urlQueryEscape("authorization_failed")) + return + } + c.Redirect(http.StatusFound, frontendRedirect+"#mcp_oauth_result=success") +} + +// Status reports whether the current user has authorized this service. +// +// Status godoc +// @Summary 查询 MCP OAuth 授权状态 +// @Description 返回当前用户对指定 MCP 服务是否已完成 OAuth 授权 +// @Tags MCP服务 +// @Produce json +// @Param id path string true "MCP 服务 ID" +// @Success 200 {object} map[string]interface{} "{authorized: bool}" +// @Security Bearer +// @Router /mcp-services/{id}/oauth/status [get] +func (h *MCPOAuthHandler) Status(c *gin.Context) { + ctx := c.Request.Context() + serviceID := c.Param("id") + tenantID := c.GetUint64(types.TenantIDContextKey.String()) + userID, _ := types.UserIDFromContext(ctx) + if tenantID == 0 || userID == "" { + c.Error(errors.NewUnauthorizedError("authentication required")) + return + } + + authorized, err := h.oauth.IsAuthorized(ctx, tenantID, userID, serviceID) + if err != nil { + c.Error(errors.NewInternalServerError("failed to query authorization status: " + err.Error())) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "data": gin.H{"authorized": authorized}}) +} + +// Revoke removes the current user's stored token and recycles connections. +// +// Revoke godoc +// @Summary 撤销 MCP OAuth 授权 +// @Description 删除当前用户对指定 MCP 服务的 OAuth 令牌 +// @Tags MCP服务 +// @Produce json +// @Param id path string true "MCP 服务 ID" +// @Success 204 +// @Security Bearer +// @Router /mcp-services/{id}/oauth/token [delete] +func (h *MCPOAuthHandler) Revoke(c *gin.Context) { + ctx := c.Request.Context() + serviceID := c.Param("id") + tenantID := c.GetUint64(types.TenantIDContextKey.String()) + userID, _ := types.UserIDFromContext(ctx) + if tenantID == 0 || userID == "" { + c.Error(errors.NewUnauthorizedError("authentication required")) + return + } + + if err := h.oauth.Revoke(ctx, tenantID, userID, serviceID); err != nil { + c.Error(errors.NewInternalServerError("failed to revoke authorization: " + err.Error())) + return + } + // Recycle any cached connections so a subsequent call re-authorizes. + _ = h.mcpManager.CloseClient(serviceID) + c.Status(http.StatusNoContent) +} diff --git a/internal/handler/mcp_service.go b/internal/handler/mcp_service.go index 7c7f64664..7b0ea7499 100644 --- a/internal/handler/mcp_service.go +++ b/internal/handler/mcp_service.go @@ -298,6 +298,23 @@ func (h *MCPServiceHandler) UpdateMCPService(c *gin.Context) { } service.AuthConfig.CustomHeaders = headers } + // auth_type and scopes are non-secret OAuth configuration; allow them + // through the main PUT so a service can be switched to/from OAuth. + if authType, ok := authConfig["auth_type"].(string); ok { + service.AuthConfig.AuthType = types.MCPAuthType(authType) + } + if scopes, ok := authConfig["scopes"].([]interface{}); ok { + list := make([]string, 0, len(scopes)) + for _, s := range scopes { + if str, ok := s.(string); ok { + list = append(list, str) + } + } + service.AuthConfig.Scopes = list + } + if metaURL, ok := authConfig["auth_server_metadata_url"].(string); ok { + service.AuthConfig.AuthServerMetadataURL = metaURL + } } if advancedConfig, ok := updateData["advanced_config"].(map[string]interface{}); ok { service.AdvancedConfig = &types.MCPAdvancedConfig{} diff --git a/internal/mcp/client.go b/internal/mcp/client.go index 465dc7f3e..c3dbc5d3f 100644 --- a/internal/mcp/client.go +++ b/internal/mcp/client.go @@ -11,6 +11,7 @@ import ( "github.com/Tencent/WeKnora/internal/logger" "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" "github.com/mark3labs/mcp-go/client" "github.com/mark3labs/mcp-go/client/transport" "github.com/mark3labs/mcp-go/mcp" @@ -49,6 +50,13 @@ type MCPClient interface { // ClientConfig represents configuration for creating an MCP client type ClientConfig struct { Service *types.MCPService + + // OAuth wiring (only used when Service.AuthConfig.AuthType == oauth). + // The token store is scoped to (TenantID, UserID, Service.ID) so each + // user connects with their own access/refresh token. + TenantID uint64 + UserID string + OAuthRepo interfaces.MCPOAuthRepository } // mcpGoClient wraps mark3labs/mcp-go client to implement our MCPClient interface @@ -92,18 +100,33 @@ func NewMCPClient(config *ClientConfig) (MCPClient, error) { } } + // Build OAuth config when this service uses the OAuth strategy. The + // client_id comes from the dynamically-registered client persisted at + // authorization time; the token store loads the invoking user's token + // and transparently refreshes it. + oauthConfig, useOAuth, err := buildOAuthConfig(config, httpClient) + if err != nil { + return nil, err + } + // Create client based on transport type var mcpClient *client.Client - var err error switch config.Service.TransportType { case types.MCPTransportSSE: if config.Service.URL == nil || *config.Service.URL == "" { return nil, fmt.Errorf("URL is required for SSE transport") } - mcpClient, err = client.NewSSEMCPClient(*config.Service.URL, - client.WithHTTPClient(httpClient), - client.WithHeaders(headers), - ) + if useOAuth { + mcpClient, err = client.NewOAuthSSEClient(*config.Service.URL, oauthConfig, + transport.WithHTTPClient(httpClient), + transport.WithHeaders(headers), + ) + } else { + mcpClient, err = client.NewSSEMCPClient(*config.Service.URL, + client.WithHTTPClient(httpClient), + client.WithHeaders(headers), + ) + } if err != nil { return nil, fmt.Errorf("failed to create SSE client: %w", err) } @@ -111,11 +134,18 @@ func NewMCPClient(config *ClientConfig) (MCPClient, error) { if config.Service.URL == nil || *config.Service.URL == "" { return nil, fmt.Errorf("URL is required for HTTP Streamable transport") } - // For HTTP streamable, we need to use transport options - mcpClient, err = client.NewStreamableHttpClient(*config.Service.URL, - transport.WithHTTPBasicClient(httpClient), - transport.WithHTTPHeaders(headers), - ) + if useOAuth { + mcpClient, err = client.NewOAuthStreamableHttpClient(*config.Service.URL, oauthConfig, + transport.WithHTTPBasicClient(httpClient), + transport.WithHTTPHeaders(headers), + ) + } else { + // For HTTP streamable, we need to use transport options + mcpClient, err = client.NewStreamableHttpClient(*config.Service.URL, + transport.WithHTTPBasicClient(httpClient), + transport.WithHTTPHeaders(headers), + ) + } if err != nil { return nil, fmt.Errorf("failed to create HTTP streamable client: %w", err) } @@ -134,6 +164,37 @@ func NewMCPClient(config *ClientConfig) (MCPClient, error) { return instance, nil } +// buildOAuthConfig returns the OAuth configuration for an OAuth-enabled MCP +// service, or (_, false, nil) when the service does not use OAuth. It loads +// the dynamically-registered client_id and wires a per-user token store so +// the transport injects the invoking user's bearer token and refreshes it. +func buildOAuthConfig(config *ClientConfig, httpClient *http.Client) (transport.OAuthConfig, bool, error) { + svc := config.Service + if !svc.AuthConfig.IsOAuth() { + return transport.OAuthConfig{}, false, nil + } + if config.OAuthRepo == nil { + return transport.OAuthConfig{}, false, fmt.Errorf("OAuth repository is required for OAuth MCP services") + } + if config.UserID == "" { + return transport.OAuthConfig{}, false, fmt.Errorf("user context is required to connect to an OAuth MCP service") + } + + oauthCfg := transport.OAuthConfig{ + Scopes: svc.AuthConfig.Scopes, + TokenStore: newDBTokenStore(config.OAuthRepo, config.TenantID, config.UserID, svc.ID), + PKCEEnabled: true, + AuthServerMetadataURL: svc.AuthConfig.AuthServerMetadataURL, + HTTPClient: httpClient, + } + if regClient, err := config.OAuthRepo.GetClient(context.Background(), config.TenantID, svc.ID); err == nil && regClient != nil { + oauthCfg.ClientID = regClient.ClientID + oauthCfg.ClientSecret = regClient.ClientSecret + oauthCfg.RedirectURI = regClient.RedirectURI + } + return oauthCfg, true, nil +} + // onConnectionLost callback when the connection is lost func (c *mcpGoClient) onConnectionLost(err error) { _ = c.Disconnect() diff --git a/internal/mcp/manager.go b/internal/mcp/manager.go index 4edb9adf7..08765d0b7 100644 --- a/internal/mcp/manager.go +++ b/internal/mcp/manager.go @@ -3,29 +3,34 @@ package mcp import ( "context" "fmt" + "strings" "sync" "time" "github.com/Tencent/WeKnora/internal/logger" "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" ) // MCPManager manages MCP client connections type MCPManager struct { - clients map[string]MCPClient // serviceID -> client + clients map[string]MCPClient // cacheKey -> client clientsMu sync.RWMutex + oauthRepo interfaces.MCPOAuthRepository ctx context.Context cancel context.CancelFunc } -// NewMCPManager creates a new MCP manager -func NewMCPManager() *MCPManager { +// NewMCPManager creates a new MCP manager. oauthRepo is used to wire per-user +// OAuth token stores for OAuth-enabled MCP services. +func NewMCPManager(oauthRepo interfaces.MCPOAuthRepository) *MCPManager { ctx, cancel := context.WithCancel(context.Background()) manager := &MCPManager{ - clients: make(map[string]MCPClient), - ctx: ctx, - cancel: cancel, + clients: make(map[string]MCPClient), + oauthRepo: oauthRepo, + ctx: ctx, + cancel: cancel, } // Start cleanup goroutine @@ -34,10 +39,23 @@ func NewMCPManager() *MCPManager { return manager } +// cacheKey computes the connection-cache key for a service. OAuth services are +// keyed per user (each user connects with their own token); all other services +// share a single connection per service ID. +func cacheKey(service *types.MCPService, userID string) string { + if service.AuthConfig.IsOAuth() { + return service.ID + "\x00" + userID + } + return service.ID +} + // GetOrCreateClient gets an existing client or creates a new one // Caches and reuses existing connections for SSE/HTTP Streamable // Note: Stdio transport is disabled for security reasons -func (m *MCPManager) GetOrCreateClient(service *types.MCPService) (MCPClient, error) { +// +// For OAuth-enabled services the connection is keyed per user (derived from +// ctx) so each user connects with their own token. +func (m *MCPManager) GetOrCreateClient(ctx context.Context, service *types.MCPService) (MCPClient, error) { // Check if service is enabled if !service.Enabled { return nil, fmt.Errorf("MCP service %s is not enabled", service.Name) @@ -48,9 +66,20 @@ func (m *MCPManager) GetOrCreateClient(service *types.MCPService) (MCPClient, er return nil, fmt.Errorf("stdio transport is disabled for security reasons; please use SSE or HTTP Streamable transport instead") } + var tenantID uint64 + var userID string + if service.AuthConfig.IsOAuth() { + tenantID, _ = types.TenantIDFromContext(ctx) + userID, _ = types.UserIDFromContext(ctx) + if userID == "" { + return nil, fmt.Errorf("user context is required to connect to OAuth MCP service %s", service.Name) + } + } + key := cacheKey(service, userID) + // For SSE/HTTP Streamable, check if client already exists and reuse m.clientsMu.RLock() - client, exists := m.clients[service.ID] + client, exists := m.clients[key] m.clientsMu.RUnlock() if exists && client.IsConnected() { @@ -62,14 +91,17 @@ func (m *MCPManager) GetOrCreateClient(service *types.MCPService) (MCPClient, er defer m.clientsMu.Unlock() // Double check after acquiring write lock - client, exists = m.clients[service.ID] + client, exists = m.clients[key] if exists && client.IsConnected() { return client, nil } // Create new client config := &ClientConfig{ - Service: service, + Service: service, + TenantID: tenantID, + UserID: userID, + OAuthRepo: m.oauthRepo, } client, err := NewMCPClient(config) @@ -89,7 +121,7 @@ func (m *MCPManager) GetOrCreateClient(service *types.MCPService) (MCPClient, er } // Store client (only for non-stdio transports) - m.clients[service.ID] = client + m.clients[key] = client logger.GetLogger(m.ctx).Infof("MCP client created and initialized for service: %s", service.Name) return client, nil @@ -128,22 +160,25 @@ func (m *MCPManager) GetClient(serviceID string) (MCPClient, bool) { return client, exists } -// CloseClient closes and removes a specific client +// CloseClient closes and removes all cached connections for a service. For +// OAuth services this spans every per-user connection (keys are prefixed with +// the service ID). func (m *MCPManager) CloseClient(serviceID string) error { m.clientsMu.Lock() defer m.clientsMu.Unlock() - client, exists := m.clients[serviceID] - if !exists { - return nil + for key, client := range m.clients { + // Match the plain service-ID key as well as per-user OAuth keys + // ("\x00"). + if key != serviceID && !strings.HasPrefix(key, serviceID+"\x00") { + continue + } + if err := client.Disconnect(); err != nil { + logger.GetLogger(m.ctx).Errorf("Failed to disconnect MCP client %s: %v", key, err) + } + delete(m.clients, key) + logger.GetLogger(m.ctx).Infof("MCP client closed: %s", key) } - - if err := client.Disconnect(); err != nil { - logger.GetLogger(m.ctx).Errorf("Failed to disconnect MCP client %s: %v", serviceID, err) - } - - delete(m.clients, serviceID) - logger.GetLogger(m.ctx).Infof("MCP client closed: %s", serviceID) return nil } diff --git a/internal/mcp/oauth_manager.go b/internal/mcp/oauth_manager.go new file mode 100644 index 000000000..14fefca39 --- /dev/null +++ b/internal/mcp/oauth_manager.go @@ -0,0 +1,189 @@ +package mcp + +import ( + "context" + "fmt" + + "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" + "github.com/mark3labs/mcp-go/client/transport" + "github.com/redis/go-redis/v9" +) + +// clientRegistrationName is sent as client_name during dynamic client +// registration (RFC 7591). +const clientRegistrationName = "WeKnora" + +// OAuthManager orchestrates the MCP OAuth2 authorization-code flow: +// discovery, dynamic client registration, the authorize redirect, and the +// callback code exchange. Tokens are persisted per (tenant, user, service); +// the registered client is persisted per (tenant, service) and reused. +type OAuthManager struct { + repo interfaces.MCPOAuthRepository + serviceRepo interfaces.MCPServiceRepository + states *oauthStateStore +} + +// NewOAuthManager constructs the OAuth manager. rdb may be nil (Lite mode), +// in which case in-flight authorization states are kept in memory. +func NewOAuthManager( + repo interfaces.MCPOAuthRepository, + serviceRepo interfaces.MCPServiceRepository, + rdb *redis.Client, +) *OAuthManager { + return &OAuthManager{ + repo: repo, + serviceRepo: serviceRepo, + states: newOAuthStateStore(rdb), + } +} + +// newHandler builds an OAuth handler bound to a service + per-user token store. +func (m *OAuthManager) newHandler( + ctx context.Context, service *types.MCPService, tenantID uint64, userID, redirectURI string, +) (*transport.OAuthHandler, error) { + if service.URL == nil || *service.URL == "" { + return nil, fmt.Errorf("MCP service URL is required for OAuth") + } + cfg := transport.OAuthConfig{ + RedirectURI: redirectURI, + Scopes: service.AuthConfig.Scopes, + TokenStore: newDBTokenStore(m.repo, tenantID, userID, service.ID), + PKCEEnabled: true, + AuthServerMetadataURL: service.AuthConfig.AuthServerMetadataURL, + } + if existing, err := m.repo.GetClient(ctx, tenantID, service.ID); err == nil && existing != nil { + cfg.ClientID = existing.ClientID + cfg.ClientSecret = existing.ClientSecret + } + h := transport.NewOAuthHandler(cfg) + h.SetBaseURL(*service.URL) + return h, nil +} + +// StartAuthorization performs discovery + (one-time) dynamic client +// registration, then returns the authorization URL the browser should visit. +// redirectURI is the backend callback URL registered with the auth server; +// frontendRedirect is where the callback bounces the browser when finished. +func (m *OAuthManager) StartAuthorization( + ctx context.Context, + service *types.MCPService, + tenantID uint64, + userID, redirectURI, frontendRedirect string, +) (string, error) { + if !service.AuthConfig.IsOAuth() { + return "", fmt.Errorf("MCP service %s does not use OAuth", service.ID) + } + + h, err := m.newHandler(ctx, service, tenantID, userID, redirectURI) + if err != nil { + return "", err + } + + // Register a client dynamically if we don't have one yet for this service. + existing, _ := m.repo.GetClient(ctx, tenantID, service.ID) + if existing == nil { + if err := h.RegisterClient(ctx, clientRegistrationName); err != nil { + return "", fmt.Errorf("dynamic client registration failed: %w", err) + } + clientID := h.GetClientID() + if clientID == "" { + return "", fmt.Errorf("dynamic client registration returned an empty client_id") + } + if err := m.repo.SaveClient(ctx, &types.MCPOAuthClient{ + TenantID: tenantID, + ServiceID: service.ID, + ClientID: clientID, + RedirectURI: redirectURI, + }); err != nil { + logger.GetLogger(ctx).Warnf("failed to persist MCP oauth client: %v", err) + } + } + + verifier, err := transport.GenerateCodeVerifier() + if err != nil { + return "", fmt.Errorf("failed to generate PKCE verifier: %w", err) + } + challenge := transport.GenerateCodeChallenge(verifier) + state, err := transport.GenerateState() + if err != nil { + return "", fmt.Errorf("failed to generate state: %w", err) + } + + authURL, err := h.GetAuthorizationURL(ctx, state, challenge) + if err != nil { + return "", fmt.Errorf("failed to build authorization URL: %w", err) + } + + if err := m.states.Put(ctx, state, OAuthState{ + TenantID: tenantID, + UserID: userID, + ServiceID: service.ID, + CodeVerifier: verifier, + ClientID: h.GetClientID(), + RedirectURI: redirectURI, + FrontendRedirect: frontendRedirect, + }); err != nil { + return "", fmt.Errorf("failed to persist authorization state: %w", err) + } + + return authURL, nil +} + +// CompleteAuthorization handles the provider callback: it validates state, +// exchanges the code for tokens (PKCE), and persists the per-user token. +// Returns the frontend redirect URL recorded at StartAuthorization time. +func (m *OAuthManager) CompleteAuthorization( + ctx context.Context, state, code string, +) (frontendRedirect string, err error) { + st, err := m.states.Take(ctx, state) + if err != nil { + return "", err + } + frontendRedirect = st.FrontendRedirect + + service, err := m.serviceRepo.GetByID(ctx, st.TenantID, st.ServiceID) + if err != nil { + return frontendRedirect, fmt.Errorf("failed to load MCP service: %w", err) + } + if service == nil { + return frontendRedirect, fmt.Errorf("MCP service not found") + } + + h, err := m.newHandler(ctx, service, st.TenantID, st.UserID, st.RedirectURI) + if err != nil { + return frontendRedirect, err + } + // Re-prime the expected state so the library's CSRF check passes after + // reconstructing the handler in this separate request. + h.SetExpectedState(state) + + if err := h.ProcessAuthorizationResponse(ctx, code, state, st.CodeVerifier); err != nil { + return frontendRedirect, fmt.Errorf("token exchange failed: %w", err) + } + // ProcessAuthorizationResponse persists the token via the TokenStore. + logger.GetLogger(ctx).Infof( + "MCP OAuth authorized: service=%s user=%s", st.ServiceID, st.UserID, + ) + return frontendRedirect, nil +} + +// IsAuthorized reports whether the given user has a stored (non-empty) token +// for the service. +func (m *OAuthManager) IsAuthorized( + ctx context.Context, tenantID uint64, userID, serviceID string, +) (bool, error) { + tok, err := m.repo.GetToken(ctx, tenantID, userID, serviceID) + if err != nil { + return false, err + } + return tok != nil && tok.AccessToken != "", nil +} + +// Revoke removes the user's stored token for the service. +func (m *OAuthManager) Revoke( + ctx context.Context, tenantID uint64, userID, serviceID string, +) error { + return m.repo.DeleteToken(ctx, tenantID, userID, serviceID) +} diff --git a/internal/mcp/oauth_state.go b/internal/mcp/oauth_state.go new file mode 100644 index 000000000..b5da49a6e --- /dev/null +++ b/internal/mcp/oauth_state.go @@ -0,0 +1,125 @@ +package mcp + +import ( + "context" + "encoding/json" + "fmt" + "os" + "strings" + "sync" + "time" + + "github.com/redis/go-redis/v9" +) + +// oauthStateTTL bounds how long an in-flight authorization may take from +// "authorize-url issued" to "callback received". +const oauthStateTTL = 10 * time.Minute + +// OAuthState is the transient data needed to complete an authorization-code +// exchange. It is keyed by the opaque OAuth `state` parameter and MUST hold +// the PKCE code_verifier, which is a secret that must never reach the +// authorization server — hence server-side storage rather than encoding it +// into the state parameter. +type OAuthState struct { + TenantID uint64 `json:"tenant_id"` + UserID string `json:"user_id"` + ServiceID string `json:"service_id"` + CodeVerifier string `json:"code_verifier"` + ClientID string `json:"client_id"` + RedirectURI string `json:"redirect_uri"` + // FrontendRedirect is where the backend callback redirects the browser + // after completing (or failing) the exchange. + FrontendRedirect string `json:"frontend_redirect"` +} + +// oauthStateStore persists in-flight OAuth states. Backed by Redis when +// available (so the callback can land on any backend replica); falls back to +// a TTL in-memory map for single-instance / Lite deployments. +type oauthStateStore struct { + rdb *redis.Client + + mu sync.Mutex + mem map[string]memStateEntry +} + +type memStateEntry struct { + value OAuthState + expiresAt time.Time +} + +func newOAuthStateStore(rdb *redis.Client) *oauthStateStore { + s := &oauthStateStore{rdb: rdb, mem: make(map[string]memStateEntry)} + if rdb == nil { + go s.gcLoop() + } + return s +} + +func (s *oauthStateStore) key(state string) string { + ns := strings.TrimSpace(os.Getenv("WEKNORA_REDIS_NAMESPACE")) + if ns != "" { + return "weknora:mcp_oauth_state:" + ns + ":" + state + } + return "weknora:mcp_oauth_state:" + state +} + +// Put stores a state with a fixed TTL. +func (s *oauthStateStore) Put(ctx context.Context, state string, value OAuthState) error { + if s.rdb != nil { + data, err := json.Marshal(value) + if err != nil { + return err + } + return s.rdb.Set(ctx, s.key(state), data, oauthStateTTL).Err() + } + s.mu.Lock() + defer s.mu.Unlock() + s.mem[state] = memStateEntry{value: value, expiresAt: time.Now().Add(oauthStateTTL)} + return nil +} + +// Take retrieves and deletes a state (single-use). Returns an error if the +// state is unknown or expired. +func (s *oauthStateStore) Take(ctx context.Context, state string) (OAuthState, error) { + if s.rdb != nil { + data, err := s.rdb.GetDel(ctx, s.key(state)).Bytes() + if err != nil { + if err == redis.Nil { + return OAuthState{}, fmt.Errorf("oauth state not found or expired") + } + return OAuthState{}, err + } + var v OAuthState + if err := json.Unmarshal(data, &v); err != nil { + return OAuthState{}, err + } + return v, nil + } + s.mu.Lock() + defer s.mu.Unlock() + entry, ok := s.mem[state] + if !ok { + return OAuthState{}, fmt.Errorf("oauth state not found or expired") + } + delete(s.mem, state) + if time.Now().After(entry.expiresAt) { + return OAuthState{}, fmt.Errorf("oauth state not found or expired") + } + return entry.value, nil +} + +func (s *oauthStateStore) gcLoop() { + ticker := time.NewTicker(time.Minute) + defer ticker.Stop() + for range ticker.C { + now := time.Now() + s.mu.Lock() + for k, v := range s.mem { + if now.After(v.expiresAt) { + delete(s.mem, k) + } + } + s.mu.Unlock() + } +} diff --git a/internal/mcp/oauth_tokenstore.go b/internal/mcp/oauth_tokenstore.go new file mode 100644 index 000000000..aa3ad02e4 --- /dev/null +++ b/internal/mcp/oauth_tokenstore.go @@ -0,0 +1,70 @@ +package mcp + +import ( + "context" + "time" + + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" + "github.com/mark3labs/mcp-go/client/transport" +) + +// dbTokenStore is a transport.TokenStore backed by the MCPOAuthRepository, +// scoped to a single (tenant, user, service) tuple. The mcp-go OAuth handler +// calls GetToken before each request (refreshing via refresh_token when +// expired) and SaveToken after a successful authorization or refresh, so this +// store transparently persists refreshed tokens back to the database. +type dbTokenStore struct { + repo interfaces.MCPOAuthRepository + tenantID uint64 + userID string + serviceID string +} + +// newDBTokenStore creates a per-user, per-service token store. +func newDBTokenStore( + repo interfaces.MCPOAuthRepository, tenantID uint64, userID, serviceID string, +) *dbTokenStore { + return &dbTokenStore{repo: repo, tenantID: tenantID, userID: userID, serviceID: serviceID} +} + +// GetToken returns the persisted token, or transport.ErrNoToken when the user +// has not authorized this service yet. +func (s *dbTokenStore) GetToken(ctx context.Context) (*transport.Token, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + row, err := s.repo.GetToken(ctx, s.tenantID, s.userID, s.serviceID) + if err != nil { + return nil, err + } + if row == nil || row.AccessToken == "" { + return nil, transport.ErrNoToken + } + return &transport.Token{ + AccessToken: row.AccessToken, + RefreshToken: row.RefreshToken, + TokenType: row.TokenType, + ExpiresAt: row.ExpiresAt, + }, nil +} + +// SaveToken persists a freshly issued or refreshed token. +func (s *dbTokenStore) SaveToken(ctx context.Context, token *transport.Token) error { + if err := ctx.Err(); err != nil { + return err + } + expiresAt := token.ExpiresAt + if expiresAt.IsZero() && token.ExpiresIn > 0 { + expiresAt = time.Now().Add(time.Duration(token.ExpiresIn) * time.Second) + } + return s.repo.SaveToken(ctx, &types.MCPOAuthToken{ + TenantID: s.tenantID, + UserID: s.userID, + ServiceID: s.serviceID, + AccessToken: token.AccessToken, + RefreshToken: token.RefreshToken, + TokenType: token.TokenType, + ExpiresAt: expiresAt, + }) +} diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go index 579246bb8..5dff59d58 100644 --- a/internal/middleware/auth.go +++ b/internal/middleware/auth.go @@ -36,6 +36,10 @@ var noAuthAPI = map[string][]string{ "/api/v1/auth/oidc/config": {"GET"}, "/api/v1/auth/oidc/url": {"GET"}, "/api/v1/auth/oidc/callback": {"GET"}, + // MCP OAuth provider redirect: the third-party authorization server + // redirects the browser here without a WeKnora bearer token. The request + // is authenticated by the opaque, single-use `state` parameter instead. + "/api/v1/mcp-oauth/callback": {"GET"}, "/api/v1/auth/refresh": {"POST"}, // IM platforms (Feishu, Slack, etc.) commonly issue a HEAD request // before GET to validate Content-Type / Content-Length when rendering diff --git a/internal/router/router.go b/internal/router/router.go index 1cf9c2f8d..2d2eb8ba9 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -67,6 +67,7 @@ type RouterParams struct { SystemHandler *handler.SystemHandler MCPServiceHandler *handler.MCPServiceHandler MCPCredentialsHandler *handler.MCPCredentialsHandler + MCPOAuthHandler *handler.MCPOAuthHandler WebSearchHandler *handler.WebSearchHandler WebSearchProviderHandler *handler.WebSearchProviderHandler WebSearchCredentialsHandler *handler.WebSearchProviderCredentialsHandler @@ -216,7 +217,7 @@ func NewRouter(params RouterParams) *gin.Engine { RegisterInitializationRoutes(v1, params.InitializationHandler, rbacGuards) RegisterSystemRoutes(v1, params.SystemHandler, rbacGuards) RegisterSystemAdminRoutes(v1, params.SystemHandler, params.AuditLogHandler, rbacGuards) - RegisterMCPServiceRoutes(v1, params.MCPServiceHandler, params.MCPCredentialsHandler, rbacGuards) + RegisterMCPServiceRoutes(v1, params.MCPServiceHandler, params.MCPCredentialsHandler, params.MCPOAuthHandler, rbacGuards) RegisterWebSearchRoutes(v1, params.WebSearchHandler, rbacGuards) RegisterWebSearchProviderRoutes(v1, params.WebSearchProviderHandler, params.WebSearchCredentialsHandler, rbacGuards) RegisterVectorStoreRoutes(v1, params.VectorStoreHandler, rbacGuards) @@ -833,8 +834,15 @@ func RegisterMCPServiceRoutes( r *gin.RouterGroup, handler *handler.MCPServiceHandler, credHandler *handler.MCPCredentialsHandler, + oauthHandler *handler.MCPOAuthHandler, g *rbacGuards, ) { + // MCP OAuth provider redirect. Registered OUTSIDE the /mcp-services group + // to avoid a static-vs-":id" route conflict, and left unauthenticated + // (allow-listed in middleware/auth.go) because the third-party browser + // redirect carries no WeKnora bearer — the single-use state authenticates. + r.GET("/mcp-oauth/callback", oauthHandler.Callback) + mcpServices := r.Group("/mcp-services") { // Create MCP service — Admin+ @@ -860,6 +868,12 @@ func RegisterMCPServiceRoutes( // MCP tool human approval (issue #1173) — Viewer+ to read, Admin+ to set policy mcpServices.GET("/:id/tool-approvals", g.Viewer(), handler.ListMCPToolApprovals) mcpServices.PUT("/:id/tool-approvals/:tool_name", g.Admin(), handler.SetMCPToolApproval) + // Per-user OAuth authorization flow. Viewer+ may authorize/inspect/ + // revoke their own token; the callback is the separate public route + // registered above. + mcpServices.POST("/:id/oauth/authorize-url", g.Viewer(), oauthHandler.AuthorizeURL) + mcpServices.GET("/:id/oauth/status", g.Viewer(), oauthHandler.Status) + mcpServices.DELETE("/:id/oauth/token", g.Viewer(), oauthHandler.Revoke) } agentTool := r.Group("/agent") diff --git a/internal/types/interfaces/mcp_oauth.go b/internal/types/interfaces/mcp_oauth.go new file mode 100644 index 000000000..cfc2a23de --- /dev/null +++ b/internal/types/interfaces/mcp_oauth.go @@ -0,0 +1,32 @@ +package interfaces + +import ( + "context" + + "github.com/Tencent/WeKnora/internal/types" +) + +// MCPOAuthRepository persists OAuth clients (per service) and tokens +// (per user + service) for the MCP OAuth2 authorization-code flow. +type MCPOAuthRepository interface { + // GetClient returns the registered OAuth client for a service, or + // (nil, nil) when none has been registered yet. + GetClient(ctx context.Context, tenantID uint64, serviceID string) (*types.MCPOAuthClient, error) + + // SaveClient creates or updates the registered OAuth client for a service. + SaveClient(ctx context.Context, client *types.MCPOAuthClient) error + + // DeleteClient removes the registered OAuth client for a service. + DeleteClient(ctx context.Context, tenantID uint64, serviceID string) error + + // GetToken returns the stored token for (tenant, user, service), or + // (nil, nil) when the user has not authorized yet. + GetToken(ctx context.Context, tenantID uint64, userID, serviceID string) (*types.MCPOAuthToken, error) + + // SaveToken creates or updates the per-user token for a service. + SaveToken(ctx context.Context, token *types.MCPOAuthToken) error + + // DeleteToken removes the per-user token for a service (revoke / + // re-authorize). + DeleteToken(ctx context.Context, tenantID uint64, userID, serviceID string) error +} diff --git a/internal/types/mcp.go b/internal/types/mcp.go index e44e84dc0..2043a2c4c 100644 --- a/internal/types/mcp.go +++ b/internal/types/mcp.go @@ -43,6 +43,22 @@ type MCPService struct { // MCPHeaders represents HTTP headers as a map type MCPHeaders map[string]string +// MCPAuthType enumerates the authentication strategies for an MCP service. +type MCPAuthType string + +const ( + // MCPAuthNone means no authentication (or only static custom headers). + MCPAuthNone MCPAuthType = "" + // MCPAuthAPIKey injects a static API key header (X-API-Key). + MCPAuthAPIKey MCPAuthType = "api_key" + // MCPAuthBearer injects a static Authorization: Bearer header. + MCPAuthBearer MCPAuthType = "bearer" + // MCPAuthOAuth performs the MCP OAuth2 authorization-code flow + // (discovery + dynamic client registration + PKCE) per user. Tokens are + // stored per (tenant, user, service) in mcp_oauth_tokens. + MCPAuthOAuth MCPAuthType = "oauth" +) + // MCPAuthConfig represents authentication configuration for MCP service. // // Secret fields (APIKey, Token) are persisted in this struct but are NEVER @@ -50,10 +66,29 @@ type MCPHeaders map[string]string // dto.MCPServiceResponse which omits them by construction. Credential // mutations happen through the dedicated /credentials subresource handled // by MCPCredentialsHandler. +// +// OAuth note: the OAuth strategy stores no secret in this struct. The +// per-user access/refresh tokens live in mcp_oauth_tokens and the +// dynamically registered client lives in mcp_oauth_clients. The fields here +// (Scopes, AuthServerMetadataURL) are non-secret OAuth configuration. type MCPAuthConfig struct { + // AuthType selects the authentication strategy. Empty ("") is treated as + // none for backward compatibility with rows that pre-date this field. + AuthType MCPAuthType `json:"auth_type,omitempty"` APIKey string `json:"api_key,omitempty"` Token string `json:"token,omitempty"` CustomHeaders map[string]string `json:"custom_headers,omitempty"` + // Scopes are the OAuth scopes requested during authorization. Optional. + Scopes []string `json:"scopes,omitempty"` + // AuthServerMetadataURL optionally pins the OAuth authorization server + // metadata URL. When empty, the server is discovered automatically from + // the MCP URL (RFC 9728 / RFC 8414). + AuthServerMetadataURL string `json:"auth_server_metadata_url,omitempty"` +} + +// IsOAuth reports whether this service uses the OAuth strategy. +func (c *MCPAuthConfig) IsOAuth() bool { + return c != nil && c.AuthType == MCPAuthOAuth } // MCPAdvancedConfig represents advanced configuration for MCP service diff --git a/internal/types/mcp_oauth.go b/internal/types/mcp_oauth.go new file mode 100644 index 000000000..8b77d8627 --- /dev/null +++ b/internal/types/mcp_oauth.go @@ -0,0 +1,139 @@ +package types + +import ( + "time" + + "github.com/Tencent/WeKnora/internal/utils" + "github.com/google/uuid" + "gorm.io/gorm" +) + +// MCPOAuthClient stores the OAuth client credentials obtained for an MCP +// service. For servers that support RFC 7591 Dynamic Client Registration the +// client_id (and optional client_secret) is registered once per service and +// reused across all users of that service, avoiding a registration round-trip +// on every authorization. +// +// One row per (tenant_id, service_id). The client_secret is encrypted at rest +// (AES-256-GCM) when SYSTEM_AES_KEY is configured. +type MCPOAuthClient struct { + ID string `json:"id" gorm:"type:varchar(36);primaryKey"` + TenantID uint64 `json:"tenant_id" gorm:"not null;uniqueIndex:idx_mcp_oauth_clients_tenant_svc"` + ServiceID string `json:"service_id" gorm:"type:varchar(36);not null;uniqueIndex:idx_mcp_oauth_clients_tenant_svc;index"` + ClientID string `json:"client_id" gorm:"type:varchar(512);not null"` + ClientSecret string `json:"-" gorm:"type:text"` + RedirectURI string `json:"redirect_uri" gorm:"type:varchar(1024)"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// TableName pins the table name. GORM's default naming would otherwise turn +// "MCPOAuthClient" into "mcpo_auth_clients" (treating "MCPO" as one token), +// which does not match the migration table "mcp_oauth_clients". +func (MCPOAuthClient) TableName() string { return "mcp_oauth_clients" } + +// BeforeCreate sets the primary key and encrypts the client secret. +func (m *MCPOAuthClient) BeforeCreate(tx *gorm.DB) error { + if m.ID == "" { + m.ID = uuid.New().String() + } + m.encryptSecret() + return nil +} + +// BeforeSave re-encrypts the secret on update paths. +func (m *MCPOAuthClient) BeforeSave(tx *gorm.DB) error { + m.encryptSecret() + return nil +} + +// AfterFind decrypts the client secret after loading. +func (m *MCPOAuthClient) AfterFind(tx *gorm.DB) error { + if plain, ok := utils.DecryptStoredSecretLenient(m.ClientSecret); ok { + m.ClientSecret = plain + } else { + m.ClientSecret = "" + } + return nil +} + +func (m *MCPOAuthClient) encryptSecret() { + if m.ClientSecret == "" { + return + } + if key := utils.GetAESKey(); key != nil { + if enc, err := utils.EncryptAESGCM(m.ClientSecret, key); err == nil { + m.ClientSecret = enc + } + } +} + +// MCPOAuthToken stores a per-user OAuth token for an MCP service. The agent +// connects to the MCP server on behalf of the invoking user, so tokens are +// isolated by (tenant_id, user_id, service_id). +// +// AccessToken and RefreshToken are encrypted at rest (AES-256-GCM) when +// SYSTEM_AES_KEY is configured. +type MCPOAuthToken struct { + ID string `json:"id" gorm:"type:varchar(36);primaryKey"` + TenantID uint64 `json:"tenant_id" gorm:"not null;uniqueIndex:idx_mcp_oauth_tokens_tenant_user_svc"` + UserID string `json:"user_id" gorm:"type:varchar(64);not null;uniqueIndex:idx_mcp_oauth_tokens_tenant_user_svc;index"` + ServiceID string `json:"service_id" gorm:"type:varchar(36);not null;uniqueIndex:idx_mcp_oauth_tokens_tenant_user_svc;index"` + AccessToken string `json:"-" gorm:"type:text"` + RefreshToken string `json:"-" gorm:"type:text"` + TokenType string `json:"token_type" gorm:"type:varchar(32)"` + ExpiresAt time.Time `json:"expires_at"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// TableName pins the table name (see MCPOAuthClient.TableName); the default +// would be "mcpo_auth_tokens" instead of the migration's "mcp_oauth_tokens". +func (MCPOAuthToken) TableName() string { return "mcp_oauth_tokens" } + +// BeforeCreate sets the primary key and encrypts secrets. +func (m *MCPOAuthToken) BeforeCreate(tx *gorm.DB) error { + if m.ID == "" { + m.ID = uuid.New().String() + } + m.encryptSecrets() + return nil +} + +// BeforeSave re-encrypts secrets on update paths. +func (m *MCPOAuthToken) BeforeSave(tx *gorm.DB) error { + m.encryptSecrets() + return nil +} + +// AfterFind decrypts secrets after loading. +func (m *MCPOAuthToken) AfterFind(tx *gorm.DB) error { + if plain, ok := utils.DecryptStoredSecretLenient(m.AccessToken); ok { + m.AccessToken = plain + } else { + m.AccessToken = "" + } + if plain, ok := utils.DecryptStoredSecretLenient(m.RefreshToken); ok { + m.RefreshToken = plain + } else { + m.RefreshToken = "" + } + return nil +} + +func (m *MCPOAuthToken) encryptSecrets() { + key := utils.GetAESKey() + if key == nil { + return + } + if m.AccessToken != "" { + if enc, err := utils.EncryptAESGCM(m.AccessToken, key); err == nil { + m.AccessToken = enc + } + } + if m.RefreshToken != "" { + if enc, err := utils.EncryptAESGCM(m.RefreshToken, key); err == nil { + m.RefreshToken = enc + } + } +} diff --git a/migrations/sqlite/000000_init.up.sql b/migrations/sqlite/000000_init.up.sql index 10cfadbe5..9a98d57ba 100644 --- a/migrations/sqlite/000000_init.up.sql +++ b/migrations/sqlite/000000_init.up.sql @@ -417,6 +417,39 @@ CREATE TABLE IF NOT EXISTS mcp_tool_approvals ( CREATE UNIQUE INDEX IF NOT EXISTS idx_mcp_tool_approvals_tenant_svc_tool ON mcp_tool_approvals(tenant_id, service_id, tool_name); CREATE INDEX IF NOT EXISTS idx_mcp_tool_approvals_service_id ON mcp_tool_approvals(service_id); +CREATE TABLE IF NOT EXISTS mcp_oauth_clients ( + id VARCHAR(36) PRIMARY KEY, + tenant_id INTEGER NOT NULL, + service_id VARCHAR(36) NOT NULL, + client_id VARCHAR(512) NOT NULL, + client_secret TEXT, + redirect_uri VARCHAR(1024), + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (service_id) REFERENCES mcp_services(id) ON DELETE CASCADE +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_mcp_oauth_clients_tenant_svc ON mcp_oauth_clients(tenant_id, service_id); +CREATE INDEX IF NOT EXISTS idx_mcp_oauth_clients_service_id ON mcp_oauth_clients(service_id); + +CREATE TABLE IF NOT EXISTS mcp_oauth_tokens ( + id VARCHAR(36) PRIMARY KEY, + tenant_id INTEGER NOT NULL, + user_id VARCHAR(64) NOT NULL, + service_id VARCHAR(36) NOT NULL, + access_token TEXT, + refresh_token TEXT, + token_type VARCHAR(32), + expires_at DATETIME, + created_at DATETIME DEFAULT CURRENT_TIMESTAMP, + updated_at DATETIME DEFAULT CURRENT_TIMESTAMP, + FOREIGN KEY (service_id) REFERENCES mcp_services(id) ON DELETE CASCADE +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_mcp_oauth_tokens_tenant_user_svc ON mcp_oauth_tokens(tenant_id, user_id, service_id); +CREATE INDEX IF NOT EXISTS idx_mcp_oauth_tokens_service_id ON mcp_oauth_tokens(service_id); +CREATE INDEX IF NOT EXISTS idx_mcp_oauth_tokens_user_id ON mcp_oauth_tokens(user_id); + CREATE TABLE IF NOT EXISTS custom_agents ( id VARCHAR(36) NOT NULL, name VARCHAR(255) NOT NULL, diff --git a/migrations/versioned/000062_mcp_oauth.down.sql b/migrations/versioned/000062_mcp_oauth.down.sql new file mode 100644 index 000000000..c860c3129 --- /dev/null +++ b/migrations/versioned/000062_mcp_oauth.down.sql @@ -0,0 +1,2 @@ +DROP TABLE IF EXISTS mcp_oauth_tokens; +DROP TABLE IF EXISTS mcp_oauth_clients; diff --git a/migrations/versioned/000062_mcp_oauth.up.sql b/migrations/versioned/000062_mcp_oauth.up.sql new file mode 100644 index 000000000..b580bd5a4 --- /dev/null +++ b/migrations/versioned/000062_mcp_oauth.up.sql @@ -0,0 +1,38 @@ +-- MCP OAuth2 support: per-service dynamically-registered clients and +-- per-user access/refresh tokens (issue: MCP OAuth2 authorization-code flow). +DO $$ BEGIN RAISE NOTICE '[Migration 000062] Creating mcp_oauth_clients / mcp_oauth_tokens...'; END $$; + +CREATE TABLE IF NOT EXISTS mcp_oauth_clients ( + id VARCHAR(36) PRIMARY KEY, + tenant_id INTEGER NOT NULL, + service_id VARCHAR(36) NOT NULL REFERENCES mcp_services(id) ON DELETE CASCADE, + client_id VARCHAR(512) NOT NULL, + client_secret TEXT, + redirect_uri VARCHAR(1024), + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_mcp_oauth_clients_tenant_svc + ON mcp_oauth_clients(tenant_id, service_id); +CREATE INDEX IF NOT EXISTS idx_mcp_oauth_clients_service_id ON mcp_oauth_clients(service_id); + +CREATE TABLE IF NOT EXISTS mcp_oauth_tokens ( + id VARCHAR(36) PRIMARY KEY, + tenant_id INTEGER NOT NULL, + user_id VARCHAR(64) NOT NULL, + service_id VARCHAR(36) NOT NULL REFERENCES mcp_services(id) ON DELETE CASCADE, + access_token TEXT, + refresh_token TEXT, + token_type VARCHAR(32), + expires_at TIMESTAMP, + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE UNIQUE INDEX IF NOT EXISTS idx_mcp_oauth_tokens_tenant_user_svc + ON mcp_oauth_tokens(tenant_id, user_id, service_id); +CREATE INDEX IF NOT EXISTS idx_mcp_oauth_tokens_service_id ON mcp_oauth_tokens(service_id); +CREATE INDEX IF NOT EXISTS idx_mcp_oauth_tokens_user_id ON mcp_oauth_tokens(user_id); + +DO $$ BEGIN RAISE NOTICE '[Migration 000062] mcp_oauth tables ready'; END $$;