mirror of
https://hubproxy.babadafafafafa.cn/https://github.com/yunionio/cloudpods.git
synced 2026-09-20 08:03:53 +08:00
feature(mcp-server): add InitAuth & userCred (#23972)
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user