feature(mcp-server): add InitAuth & userCred (#23972)

This commit is contained in:
cwz_eikoh
2025-12-30 10:36:45 +08:00
committed by GitHub
parent 7dfe635b62
commit 36d6188385
10 changed files with 105 additions and 42 deletions

View File

@@ -17,7 +17,11 @@ package adapters
import (
"context"
"yunion.io/x/log"
api "yunion.io/x/onecloud/pkg/apis/identity"
"yunion.io/x/onecloud/pkg/cloudcommon/policy"
"yunion.io/x/onecloud/pkg/mcclient"
"yunion.io/x/onecloud/pkg/mcclient/auth"
"yunion.io/x/onecloud/pkg/mcp-server/options"
)
@@ -35,7 +39,7 @@ type CloudRegion struct {
func NewCloudpodsAdapter() *CloudpodsAdapter {
client := mcclient.NewClient(
options.Options.IdentityBaseURL,
options.Options.AuthURL,
options.Options.Timeout,
false,
true,
@@ -49,30 +53,46 @@ func NewCloudpodsAdapter() *CloudpodsAdapter {
}
// authenticate 实现 Cloudpods 的认证逻辑,例如获取访问令牌
func (a *CloudpodsAdapter) authenticate(ak string, sk string) error {
func (a *CloudpodsAdapter) authenticate(ak string, sk string) (mcclient.TokenCredential, error) {
if a.session != nil {
return nil
return a.session.GetToken(), nil
}
token, err := a.client.AuthenticateByAccessKey(ak, sk, "")
if err != nil {
return err
return nil, err
}
a.session = a.client.NewSession(
context.Background(),
"",
"",
"apigateway",
token,
)
return nil
return token, nil
}
func (a *CloudpodsAdapter) getSession(ak string, sk string) (*mcclient.ClientSession, error) {
if err := a.authenticate(ak, sk); err != nil {
return nil, err
func (a *CloudpodsAdapter) getSession(ctx context.Context, ak string, sk string) (*mcclient.ClientSession, error) {
var userCred mcclient.TokenCredential
if auth.IsAuthed() {
userCred = policy.FetchUserCredential(ctx)
if userCred != nil {
log.Infof("getSessionWithUserCred: %v", userCred)
} else {
log.Infof("No userCred in context, will use ak/sk for authentication")
token, err := a.authenticate(ak, sk)
if err != nil {
return nil, err
}
userCred = token
}
a.session = auth.GetSession(ctx, userCred, "")
} else {
token, err := a.authenticate(ak, sk)
if err != nil {
return nil, err
}
a.session = a.client.NewSession(
context.Background(),
"",
"",
api.EndpointInterfaceApigateway,
token,
)
}
return a.session, nil
}

View File

@@ -31,7 +31,7 @@ import (
// StartServer 启动 Cloudpods 中的服务器
func (a *CloudpodsAdapter) StartServer(ctx context.Context, serverId string, req models.ServerStartRequest, ak string, sk string) (*models.ServerOperationResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -75,7 +75,7 @@ func (a *CloudpodsAdapter) StartServer(ctx context.Context, serverId string, req
// StopServer 停止 Cloudpods 中的服务器
func (a *CloudpodsAdapter) StopServer(ctx context.Context, serverId string, req models.ServerStopRequest, ak string, sk string) (*models.ServerOperationResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -124,7 +124,7 @@ func (a *CloudpodsAdapter) StopServer(ctx context.Context, serverId string, req
// RestartServer 重启 Cloudpods 中的服务器
func (a *CloudpodsAdapter) RestartServer(ctx context.Context, serverId string, req models.ServerRestartRequest, ak string, sk string) (*models.ServerOperationResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -163,7 +163,7 @@ func (a *CloudpodsAdapter) RestartServer(ctx context.Context, serverId string, r
// ResetServerPassword 重置 Cloudpods 中服务器的密码
func (a *CloudpodsAdapter) ResetServerPassword(ctx context.Context, serverId string, req models.ServerResetPasswordRequest, ak string, sk string) (*models.ServerOperationResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -211,7 +211,7 @@ func (a *CloudpodsAdapter) ResetServerPassword(ctx context.Context, serverId str
// DeleteServer 删除 Cloudpods 中的服务器
func (a *CloudpodsAdapter) DeleteServer(ctx context.Context, serverId string, req models.ServerDeleteRequest, ak string, sk string) (*models.ServerOperationResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -265,7 +265,7 @@ func (a *CloudpodsAdapter) DeleteServer(ctx context.Context, serverId string, re
// CreateServer 在 Cloudpods 中创建服务器
func (a *CloudpodsAdapter) CreateServer(ctx context.Context, req models.CreateServerRequest, ak string, sk string) (*models.CreateServerResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -440,7 +440,7 @@ func (a *CloudpodsAdapter) CreateServer(ctx context.Context, req models.CreateSe
// GetServerMonitor 获取 Cloudpods 中服务器的监控数据
func (a *CloudpodsAdapter) GetServerMonitor(ctx context.Context, serverId string, startTime, endTime int64, metrics []string, ak string, sk string) (*models.MonitorResponse, error) {
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -591,7 +591,7 @@ func (a *CloudpodsAdapter) GetServerMonitor(ctx context.Context, serverId string
// GetServerStats 获取 Cloudpods 中服务器的实时统计数据
func (a *CloudpodsAdapter) GetServerStats(ctx context.Context, serverId string, ak string, sk string) (*models.ServerStatsResponse, error) {
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}

View File

@@ -29,7 +29,7 @@ import (
// ListCloudRegions 查询 Cloudpods 中的区域列表
func (a CloudpodsAdapter) ListCloudRegions(ctx context.Context, limit int, offset int, search string, provider string, ak string, sk string) (*models.CloudregionListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -83,9 +83,9 @@ func (a CloudpodsAdapter) ListCloudRegions(ctx context.Context, limit int, offse
}
// ListVPCs 查询 Cloudpods 中的 VPC 列表
func (a *CloudpodsAdapter) ListVPCs(limit int, offset int, search string, cloudregionId string, ak string, sk string) (*models.VpcListResponse, error) {
func (a *CloudpodsAdapter) ListVPCs(ctx context.Context, limit int, offset int, search string, cloudregionId string, ak string, sk string) (*models.VpcListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -139,9 +139,9 @@ func (a *CloudpodsAdapter) ListVPCs(limit int, offset int, search string, cloudr
}
// ListNetworks 查询 Cloudpods 中的网络列表
func (a *CloudpodsAdapter) ListNetworks(limit int, offset int, search string, vpcId string, ak string, sk string) (*models.NetworkListResponse, error) {
func (a *CloudpodsAdapter) ListNetworks(ctx context.Context, limit int, offset int, search string, vpcId string, ak string, sk string) (*models.NetworkListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -196,9 +196,9 @@ func (a *CloudpodsAdapter) ListNetworks(limit int, offset int, search string, vp
}
// ListImages 查询 Cloudpods 中的镜像列表
func (a *CloudpodsAdapter) ListImages(limit int, offset int, search string, osTypes []string, ak string, sk string) (*models.ImageListResponse, error) {
func (a *CloudpodsAdapter) ListImages(ctx context.Context, limit int, offset int, search string, osTypes []string, ak string, sk string) (*models.ImageListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -254,9 +254,9 @@ func (a *CloudpodsAdapter) ListImages(limit int, offset int, search string, osTy
}
// ListServerSkus 查询 Cloudpods 中的服务器规格列表
func (a *CloudpodsAdapter) ListServerSkus(limit int, offset int, search string, cloudregionIds []string, zoneIds []string, cpuCoreCount []string, memorySizeMB []string, providers []string, cpuArch []string, ak string, sk string) (*models.ServerSkuListResponse, error) {
func (a *CloudpodsAdapter) ListServerSkus(ctx context.Context, limit int, offset int, search string, cloudregionIds []string, zoneIds []string, cpuCoreCount []string, memorySizeMB []string, providers []string, cpuArch []string, ak string, sk string) (*models.ServerSkuListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -352,9 +352,9 @@ func (a *CloudpodsAdapter) ListServerSkus(limit int, offset int, search string,
}
// ListStorages 查询 Cloudpods 中的存储列表
func (a *CloudpodsAdapter) ListStorages(limit int, offset int, search string, cloudregionIds []string, zoneIds []string, providers []string, storageTypes []string, hostId string, ak string, sk string) (*models.StorageListResponse, error) {
func (a *CloudpodsAdapter) ListStorages(ctx context.Context, limit int, offset int, search string, cloudregionIds []string, zoneIds []string, providers []string, storageTypes []string, hostId string, ak string, sk string) (*models.StorageListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}
@@ -440,7 +440,7 @@ func (a *CloudpodsAdapter) ListStorages(limit int, offset int, search string, cl
// ListServers 查询 Cloudpods 中的服务器列表
func (a *CloudpodsAdapter) ListServers(ctx context.Context, limit int, offset int, search string, status string, ak string, sk string) (*models.ServerListResponse, error) {
// 获取 Cloudpods 会话
session, err := a.getSession(ak, sk)
session, err := a.getSession(ctx, ak, sk)
if err != nil {
return nil, err
}

View File

@@ -15,11 +15,17 @@
package server
import (
"context"
"fmt"
"net/http"
"github.com/mark3labs/mcp-go/server"
"yunion.io/x/log"
"yunion.io/x/pkg/appctx"
api "yunion.io/x/onecloud/pkg/apis/identity"
"yunion.io/x/onecloud/pkg/mcclient/auth"
"yunion.io/x/onecloud/pkg/mcp-server/adapters"
"yunion.io/x/onecloud/pkg/mcp-server/options"
@@ -136,8 +142,32 @@ func (s *CloudpodsMCPServer) registerAllTools() error {
// Start 以sse模式启动 mcp 服务
func (s *CloudpodsMCPServer) Start() error {
// 设置 contextFunc 来从 HTTP header 中提取认证信息并放入 context
contextFunc := func(ctx context.Context, r *http.Request) context.Context {
tokenStr := r.Header.Get(api.AUTH_TOKEN_HEADER)
if len(tokenStr) > 0 {
if auth.IsAuthed() {
userCred, err := auth.Verify(ctx, tokenStr)
if err != nil {
log.Errorf("Verify token failed: %s", err)
return ctx
}
// 将 userCred 放入 context
ctx = context.WithValue(ctx, appctx.APP_CONTEXT_KEY_AUTH_TOKEN, userCred)
log.Debugf("UserCred set in context from HTTP header token")
} else {
log.Warningf("Auth manager not initialized, skipping token verification")
}
}
return ctx
}
if err := server.NewSSEServer(s.mcpServer).Start(fmt.Sprintf("%s:%d", options.Options.Address, options.Options.Port)); err != nil {
sseServer := server.NewSSEServer(
s.mcpServer,
server.WithSSEContextFunc(contextFunc),
)
if err := sseServer.Start(fmt.Sprintf("%s:%d", options.Options.Address, options.Options.Port)); err != nil {
return err
}
log.Infof("Start mcp server successfully")

View File

@@ -19,6 +19,7 @@ import (
"yunion.io/x/log"
app_common "yunion.io/x/onecloud/pkg/cloudcommon/app"
common_options "yunion.io/x/onecloud/pkg/cloudcommon/options"
"yunion.io/x/onecloud/pkg/mcp-server/options"
"yunion.io/x/onecloud/pkg/mcp-server/server"
@@ -29,6 +30,18 @@ func StartService() {
opts := &options.Options
common_options.ParseOptions(opts, os.Args, "mcpserver.conf", "mcpserver")
// 如果配置了认证信息,初始化 auth manager
commonOpts := &opts.CommonOptions
// 只有当所有必需的认证配置都存在时,才初始化 auth manager
if len(commonOpts.AuthURL) > 0 && len(commonOpts.AdminUser) > 0 &&
len(commonOpts.AdminPassword) > 0 && len(commonOpts.AdminProject) > 0 {
app_common.InitAuth(commonOpts, func() {
log.Infof("Auth complete!!")
})
} else {
log.Infof("Auth configuration incomplete, skipping auth initialization. AuthURL: %s, AdminUser: %s, AdminPassword: %s, AdminProject: %s", commonOpts.AuthURL, commonOpts.AdminUser, commonOpts.AdminPassword, commonOpts.AdminProject)
}
// 创建服务器
srv := server.NewServer()

View File

@@ -110,7 +110,7 @@ func (c *CloudpodsImagesTool) Handle(ctx context.Context, req mcp.CallToolReques
sk := req.GetString("sk", "")
// 调用适配器查询镜像列表
imagesResponse, err := c.adapter.ListImages(limit, offset, search, osTypes, ak, sk)
imagesResponse, err := c.adapter.ListImages(ctx, limit, offset, search, osTypes, ak, sk)
if err != nil {
log.Errorf("Fail to query image: %s", err)
return nil, fmt.Errorf("fail to query image: %w", err)

View File

@@ -97,7 +97,7 @@ func (c *CloudpodsNetworksTool) Handle(ctx context.Context, req mcp.CallToolRequ
sk := req.GetString("sk", "")
// 调用适配器获取网络列表
networksResponse, err := c.adapter.ListNetworks(limit, offset, search, vpcId, ak, sk)
networksResponse, err := c.adapter.ListNetworks(ctx, limit, offset, search, vpcId, ak, sk)
if err != nil {
log.Errorf("Fail to query network: %s", err)
return nil, fmt.Errorf("fail to query network: %w", err)

View File

@@ -171,7 +171,7 @@ func (c *CloudpodsServerSkusTool) Handle(ctx context.Context, req mcp.CallToolRe
sk := req.GetString("sk", "")
// 调用适配器查询主机套餐规格列表
skusResponse, err := c.adapter.ListServerSkus(limit, offset, search, cloudregionIds, zoneIds, cpuCoreCount, memorySizeMB, providers, cpuArch, ak, sk)
skusResponse, err := c.adapter.ListServerSkus(ctx, limit, offset, search, cloudregionIds, zoneIds, cpuCoreCount, memorySizeMB, providers, cpuArch, ak, sk)
if err != nil {
log.Errorf("Fail to query server skus: %s", err)
return nil, fmt.Errorf("fail to query server skus: %w", err)

View File

@@ -155,7 +155,7 @@ func (c *CloudpodsStoragesTool) Handle(ctx context.Context, req mcp.CallToolRequ
sk := req.GetString("sk", "")
// 调用适配器查询块存储列表
storagesResponse, err := c.adapter.ListStorages(limit, offset, search, cloudregionIds, zoneIds, providers, storageTypes, hostId, ak, sk)
storagesResponse, err := c.adapter.ListStorages(ctx, limit, offset, search, cloudregionIds, zoneIds, providers, storageTypes, hostId, ak, sk)
if err != nil {
log.Errorf("Fail to query storage: %s", err)
return nil, fmt.Errorf("fail to query storage: %w", err)

View File

@@ -111,7 +111,7 @@ func (c *CloudpodsVPCsTool) Handle(ctx context.Context, req mcp.CallToolRequest)
sk := req.GetString("sk", "")
// 调用适配器查询VPC列表
vpcsResponse, err := c.adapter.ListVPCs(limit, offset, search, cloudRegionID, ak, sk)
vpcsResponse, err := c.adapter.ListVPCs(ctx, limit, offset, search, cloudRegionID, ak, sk)
if err != nil {
log.Errorf("Fail to query vpc: %s", err)
return nil, fmt.Errorf("fail to query vpc: %w", err)