You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
159 lines
5.3 KiB
159 lines
5.3 KiB
package mcp
|
|
|
|
import (
|
|
"context"
|
|
"yunyan/comm"
|
|
"yunyan/lego/core"
|
|
"yunyan/lego/core/cbase"
|
|
"yunyan/lego/sys/log"
|
|
"yunyan/lego/sys/mysql"
|
|
"yunyan/pb"
|
|
"yunyan/utils"
|
|
"fmt"
|
|
"time"
|
|
|
|
"github.com/mark3labs/mcp-go/mcp"
|
|
)
|
|
|
|
// 添加用户任务工具
|
|
type tool_allhelp_task struct {
|
|
cbase.ModuleCompBase
|
|
module *Mcp
|
|
}
|
|
|
|
// 组件初始化接口
|
|
func (this *tool_allhelp_task) Init(service core.IService, module core.IModule, comp core.IModuleComp, opt core.IModuleOptions) (err error) {
|
|
this.ModuleCompBase.Init(service, module, comp, opt)
|
|
this.module = module.(*Mcp)
|
|
return
|
|
}
|
|
|
|
func (this *tool_allhelp_task) Start() (err error) {
|
|
err = this.ModuleCompBase.Start()
|
|
|
|
return
|
|
}
|
|
|
|
func (this *tool_allhelp_task) Tool() (tool mcp.Tool) {
|
|
return mcp.NewTool("add_user_task",
|
|
mcp.WithDescription("根据用户的需求添加一个定时提醒任务,例如:帮我设一个明天下午3点的会议提醒"),
|
|
// uid 不再是必填:身份由请求头里的 JWT 决定。保留这个参数只为兼容
|
|
// 百炼后台已配好的工具定义,传上来也只会被拿去和会话 uid 比对。
|
|
mcp.WithString("uid",
|
|
mcp.Description("用户ID(可省略,服务端以登录身份为准)"),
|
|
),
|
|
mcp.WithString("task_name",
|
|
mcp.Description("任务标题(简短摘要,如:下午3点会议提醒)"),
|
|
mcp.Required(),
|
|
),
|
|
mcp.WithString("task_description",
|
|
mcp.Description("任务详细描述"),
|
|
),
|
|
mcp.WithNumber("task_type",
|
|
mcp.Description("任务类型: 0=一次性提醒, 1=每天重复, 2=每周重复, 3=每月重复"),
|
|
mcp.DefaultNumber(0),
|
|
),
|
|
mcp.WithString("trigger_time",
|
|
mcp.Description("触发时间,必须为带时区的 ISO 8601 字符串(RFC3339),如 2025-04-08T15:00:00+08:00 或 2025-04-08T07:00:00Z。禁止使用 Unix 时间戳,禁止省略时区"),
|
|
mcp.Required(),
|
|
),
|
|
mcp.WithString("extra",
|
|
mcp.Description("可选的扩展信息JSON,如会议链接、地点等"),
|
|
),
|
|
)
|
|
}
|
|
|
|
func (this *tool_allhelp_task) Handl(ctx context.Context, request mcp.CallToolRequest) (result *mcp.CallToolResult, err error) {
|
|
this.module.Debug("tool_allhelp_task",
|
|
log.Field{Key: "request.Params.Arguments", Value: request.GetRawArguments()},
|
|
)
|
|
|
|
// ⚠️ 参数里的 uid 只用于比对,**绝不作为身份来源**。真正的身份从
|
|
// Authorization 里的 JWT 解出来(见 auth.go)。此前这里直接信任模型填的 uid,
|
|
// 而 MCP 是公网可达且无鉴权的独立服务,等于谁都能读改任意人的数据。
|
|
argUID := request.GetString("uid", "")
|
|
uid, err := this.module.ResolveUIDWithToken(ctx, argUID, request.GetString("auth_token", ""))
|
|
if err != nil {
|
|
return mcp.NewToolResultError(err.Error()), nil
|
|
}
|
|
user := &pb.DBUser{}
|
|
if findErr := mysql.FindOne(comm.TableUser, user, "uid=?", uid); findErr != nil {
|
|
return mcp.NewToolResultError("用户不存在"), nil
|
|
}
|
|
taskName, err := request.RequireString("task_name")
|
|
if err != nil {
|
|
return mcp.NewToolResultError("task_name is required"), nil
|
|
}
|
|
triggerTimeStr, err := request.RequireString("trigger_time")
|
|
if err != nil {
|
|
return mcp.NewToolResultError("trigger_time is required"), nil
|
|
}
|
|
|
|
taskDesc := request.GetString("task_description", "")
|
|
taskType := request.GetInt("task_type", 0)
|
|
extra := request.GetString("extra", "")
|
|
|
|
// 校验触发时间格式(必须是带时区的 RFC3339),原字符串按用户输入原样入库
|
|
if _, parseErr := parseTriggerTime(triggerTimeStr); parseErr != nil {
|
|
return mcp.NewToolResultError(fmt.Sprintf("无法解析触发时间: %v", parseErr)), nil
|
|
}
|
|
|
|
now := time.Now().Unix()
|
|
task := &pb.DBTask{
|
|
Uid: uid,
|
|
TaskName: taskName,
|
|
TaskDesc: taskDesc,
|
|
TaskType: pb.TaskType(taskType),
|
|
TriggerTime: triggerTimeStr,
|
|
Extra: extra,
|
|
CreateTime: now,
|
|
UpdateTime: now,
|
|
Status: pb.TaskStatus_TaskStatus_Pending,
|
|
}
|
|
|
|
if err := mysql.Insert(comm.TableAllhelpTask, task); err != nil {
|
|
this.module.Errorln("add_user_task insert error:", err)
|
|
return mcp.NewToolResultError("创建任务失败"), nil
|
|
}
|
|
|
|
this.module.Debug("tool_allhelp_task success",
|
|
log.Field{Key: "task_id", Value: task.Id},
|
|
log.Field{Key: "uid", Value: uid},
|
|
log.Field{Key: "task_name", Value: taskName},
|
|
)
|
|
|
|
text := utils.ToString(map[string]interface{}{
|
|
"result": "success",
|
|
"task_id": task.Id,
|
|
"task_name": taskName,
|
|
"task_description": taskDesc,
|
|
"task_type": taskType,
|
|
"trigger_time": triggerTimeStr,
|
|
"extra": extra,
|
|
"status": int32(task.Status),
|
|
"create_time": task.CreateTime,
|
|
"update_time": task.UpdateTime,
|
|
"message": fmt.Sprintf("已成功创建提醒任务「%s」,将在 %s 提醒您", taskName, triggerTimeStr),
|
|
})
|
|
|
|
result = &mcp.CallToolResult{
|
|
Content: []mcp.Content{
|
|
mcp.TextContent{
|
|
Type: "text",
|
|
Text: text,
|
|
},
|
|
},
|
|
}
|
|
return
|
|
}
|
|
|
|
// parseTriggerTime 解析触发时间,仅接受带时区的 ISO 8601 字符串(RFC3339)
|
|
// 禁止 Unix 时间戳与无时区格式,避免跨时区语义歧义
|
|
func parseTriggerTime(s string) (time.Time, error) {
|
|
for _, f := range []string{time.RFC3339Nano, time.RFC3339} {
|
|
if t, err := time.Parse(f, s); err == nil {
|
|
return t, nil
|
|
}
|
|
}
|
|
return time.Time{}, fmt.Errorf("trigger_time 必须为带时区的 ISO 8601 字符串(如 2025-04-08T15:00:00+08:00),收到: %s", s)
|
|
}
|
|
|