f3a44a901e
Build and Deploy / build-and-push (push) Successful in 54s
- Move MCP tools registration to internal/mcp - cmd/mcp now only contains CLI entry logic - Support env vars TASKPOOL_URL/TASKPOOL_TOKEN for stdio mode - Remove duplicate client.go from cmd/mcp
211 lines
4.9 KiB
Go
211 lines
4.9 KiB
Go
package mcp
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
|
|
"github.com/engigu/taskpool/internal/constant"
|
|
"github.com/engigu/taskpool/internal/models/vo"
|
|
"github.com/engigu/taskpool/internal/services"
|
|
)
|
|
|
|
var (
|
|
openAPIClient *Client
|
|
openAPIClientOnce sync.Once
|
|
)
|
|
|
|
// GetOpenAPIClient 获取 OpenAPI 客户端(单例)
|
|
// 优先使用环境变量 TASKPOOL_URL 和 TASKPOOL_TOKEN
|
|
// 否则使用本地服务器地址和系统设置中的 OpenAPI Token
|
|
func GetOpenAPIClient() *Client {
|
|
openAPIClientOnce.Do(func() {
|
|
// 优先从环境变量读取(支持 stdio 模式连接远程服务器)
|
|
envURL := strings.TrimSpace(os.Getenv("TASKPOOL_URL"))
|
|
envToken := strings.TrimSpace(os.Getenv("TASKPOOL_TOKEN"))
|
|
if envURL == "" {
|
|
envURL = strings.TrimSpace(os.Getenv("BAIHU_URL"))
|
|
}
|
|
if envToken == "" {
|
|
envToken = strings.TrimSpace(os.Getenv("BAIHU_TOKEN"))
|
|
}
|
|
|
|
var baseURL, token string
|
|
if envURL != "" && envToken != "" {
|
|
// 使用环境变量配置
|
|
baseURL = envURL
|
|
token = envToken
|
|
} else {
|
|
// 使用本地服务器配置
|
|
cfg := services.GetConfig()
|
|
baseURL = fmt.Sprintf("http://127.0.0.1:%d", cfg.Server.Port)
|
|
if cfg.Server.URLPrefix != "" {
|
|
baseURL += cfg.Server.URLPrefix
|
|
}
|
|
settingsSvc := services.NewSettingsService()
|
|
token = getOpenAPIToken(settingsSvc)
|
|
}
|
|
|
|
openAPIClient = NewClient(baseURL, token)
|
|
})
|
|
return openAPIClient
|
|
}
|
|
|
|
// RefreshClient 刷新客户端(当 OpenAPI Token 变更时调用)
|
|
func RefreshClient() {
|
|
settingsSvc := services.NewSettingsService()
|
|
token := getOpenAPIToken(settingsSvc)
|
|
if openAPIClient != nil {
|
|
openAPIClient.Token = token
|
|
}
|
|
}
|
|
|
|
func getOpenAPIToken(settingsSvc *services.SettingsService) string {
|
|
siteConfig := settingsSvc.GetSection(constant.SectionSite)
|
|
tokenJson, ok := siteConfig[constant.KeyOpenapiToken]
|
|
if !ok || tokenJson == "" {
|
|
return ""
|
|
}
|
|
|
|
var tokenConfig vo.TokenConfig
|
|
if err := json.Unmarshal([]byte(tokenJson), &tokenConfig); err != nil {
|
|
return ""
|
|
}
|
|
|
|
if !tokenConfig.Enabled {
|
|
return ""
|
|
}
|
|
|
|
return tokenConfig.Token
|
|
}
|
|
|
|
// Client 调用内部 OpenAPI
|
|
type Client struct {
|
|
BaseURL string
|
|
Token string
|
|
HTTPClient *http.Client
|
|
}
|
|
|
|
func NewClient(baseURL, token string) *Client {
|
|
return &Client{
|
|
BaseURL: strings.TrimRight(strings.TrimSpace(baseURL), "/"),
|
|
Token: strings.TrimSpace(token),
|
|
HTTPClient: &http.Client{
|
|
Timeout: 60 * time.Second,
|
|
},
|
|
}
|
|
}
|
|
|
|
func (c *Client) openBase() string {
|
|
return c.BaseURL + "/open2api/v1"
|
|
}
|
|
|
|
func (c *Client) do(method, path string, query map[string]string, body any) (json.RawMessage, error) {
|
|
if c.BaseURL == "" {
|
|
return nil, fmt.Errorf("服务器地址未配置")
|
|
}
|
|
if c.Token == "" {
|
|
return nil, fmt.Errorf("OpenAPI Token 未配置,请在系统设置中启用 OpenAPI 并生成 Token")
|
|
}
|
|
|
|
u, err := url.Parse(c.openBase() + path)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if len(query) > 0 {
|
|
q := u.Query()
|
|
for k, v := range query {
|
|
if strings.TrimSpace(v) != "" {
|
|
q.Set(k, v)
|
|
}
|
|
}
|
|
u.RawQuery = q.Encode()
|
|
}
|
|
|
|
var reader io.Reader
|
|
if body != nil {
|
|
b, err := json.Marshal(body)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
reader = bytes.NewReader(b)
|
|
}
|
|
|
|
req, err := http.NewRequest(method, u.String(), reader)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
req.Header.Set("Authorization", "Bearer "+c.Token)
|
|
req.Header.Set("Accept", "application/json")
|
|
if body != nil {
|
|
req.Header.Set("Content-Type", "application/json")
|
|
}
|
|
|
|
resp, err := c.HTTPClient.Do(req)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("请求失败: %w", err)
|
|
}
|
|
defer resp.Body.Close()
|
|
|
|
raw, err := io.ReadAll(resp.Body)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
var envelope struct {
|
|
Code int `json:"code"`
|
|
Msg string `json:"msg"`
|
|
Data json.RawMessage `json:"data"`
|
|
}
|
|
if err := json.Unmarshal(raw, &envelope); err != nil {
|
|
if resp.StatusCode >= 400 {
|
|
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, truncate(string(raw), 500))
|
|
}
|
|
return raw, nil
|
|
}
|
|
|
|
if envelope.Code != 200 {
|
|
msg := envelope.Msg
|
|
if msg == "" {
|
|
msg = truncate(string(raw), 500)
|
|
}
|
|
return nil, fmt.Errorf("[%d] %s", envelope.Code, msg)
|
|
}
|
|
if len(envelope.Data) > 0 && string(envelope.Data) != "null" {
|
|
return envelope.Data, nil
|
|
}
|
|
result := map[string]any{"code": envelope.Code, "msg": envelope.Msg}
|
|
b, _ := json.Marshal(result)
|
|
return b, nil
|
|
}
|
|
|
|
func (c *Client) Get(path string, query map[string]string) (json.RawMessage, error) {
|
|
return c.do(http.MethodGet, path, query, nil)
|
|
}
|
|
|
|
func (c *Client) Post(path string, body any) (json.RawMessage, error) {
|
|
return c.do(http.MethodPost, path, nil, body)
|
|
}
|
|
|
|
func (c *Client) Put(path string, body any) (json.RawMessage, error) {
|
|
return c.do(http.MethodPut, path, nil, body)
|
|
}
|
|
|
|
func (c *Client) Delete(path string) (json.RawMessage, error) {
|
|
return c.do(http.MethodDelete, path, nil, nil)
|
|
}
|
|
|
|
func truncate(s string, n int) string {
|
|
if len(s) <= n {
|
|
return s
|
|
}
|
|
return s[:n] + "..."
|
|
}
|