package mcp import ( "context" "strings" "time" "yunyan/comm" "yunyan/lego/core" "yunyan/lego/core/cbase" "yunyan/lego/sys/mysql" "yunyan/pb" "yunyan/utils" "github.com/mark3labs/mcp-go/mcp" ) /* 记忆中心(拾忆)的 MCP 工具:让 EMAI 能回答「我明天有什么事」「这个月花了多少」。 ## 只读 新增/修改/删除**不放在这里**。写路径的校验、幂等(client_key)、提醒时刻重算 都在 home 的 memory 模块内,在 MCP 里复制一份必然漂移。 录入走端侧指令(tool_calls),查询走 MCP —— 这个分工是刻意的。 ## 身份 uid 一律从 Authorization 的 JWT 解(见 auth.go),**不信任模型传的参数**。 */ // ---------- get_memory_items ---------- type tool_get_memory_items struct { cbase.ModuleCompBase module *Mcp } func (this *tool_get_memory_items) 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_get_memory_items) Start() (err error) { err = this.ModuleCompBase.Start() this.module.AddTool(ToolGroup_GLOBAL, this.Tool(), this.Handl) return } func (this *tool_get_memory_items) Tool() mcp.Tool { return mcp.NewTool("get_memory_items", mcp.WithDescription("查询用户记录的事项(待办/闹钟/灵感/花销),可按时间范围与分类筛选"), mcp.WithString("range", mcp.Description("相对时间范围,取值:"+strings.Join(comm.MemoryRelativeRanges, "/")+ "。与 start_date/end_date 二选一,优先用本字段"), ), mcp.WithString("start_date", mcp.Description("起始日期 YYYY-MM-DD")), mcp.WithString("end_date", mcp.Description("结束日期 YYYY-MM-DD")), mcp.WithString("category", mcp.Description("分类筛选:todo=待办 alarm=闹钟 idea=灵感 expense=花销,留空=全部"), ), mcp.WithNumber("limit", mcp.Description("返回条数上限(默认50,最大200)"), mcp.DefaultNumber(50)), ) } func (this *tool_get_memory_items) Handl(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { uid, err := this.module.ResolveUID(ctx, request.GetString("uid", "")) if err != nil { return mcp.NewToolResultError(err.Error()), nil } start, end, errRange := resolveRange(request) if errRange != nil { return mcp.NewToolResultError(errRange.Error()), nil } var cats []string if c := strings.TrimSpace(request.GetString("category", "")); c != "" { // 模型填错分类不该让整次查询失败——忽略非法值当作"不限分类", // 比返回一个错误让它重试要好(它多半会再填错一次)。 cats, _ = comm.NormalizeMemoryCategories([]string{c}) } limit := int(request.GetInt("limit", 50)) if limit <= 0 || limit > 200 { limit = 50 } items := make([]*pb.DBMemoryItem, 0) tx := comm.ApplyMemoryQuery(mysql.Table(comm.TableMemoryItem), comm.MemoryQuery{ Uid: uid, StartDate: start, EndDate: end, Categories: cats, }) if e := comm.OrderMemoryItems(tx).Limit(limit).Find(&items).Error; e != nil { this.module.Errorln("get_memory_items error:", e) return mcp.NewToolResultError("查询失败"), nil } return mcp.NewToolResultText(utils.ToString(map[string]interface{}{ "result": "success", "start_date": start, "end_date": end, "total": len(items), "items": briefItems(items), })), nil } // ---------- get_memory_stats ---------- type tool_get_memory_stats struct { cbase.ModuleCompBase module *Mcp } func (this *tool_get_memory_stats) 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_get_memory_stats) Start() (err error) { err = this.ModuleCompBase.Start() this.module.AddTool(ToolGroup_GLOBAL, this.Tool(), this.Handl) return } func (this *tool_get_memory_stats) Tool() mcp.Tool { return mcp.NewTool("get_memory_stats", mcp.WithDescription("统计用户某段时间的待办完成情况、灵感数量与花销合计"), mcp.WithString("range", mcp.Description("相对时间范围,取值:"+strings.Join(comm.MemoryRelativeRanges, "/")), ), mcp.WithString("start_date", mcp.Description("起始日期 YYYY-MM-DD")), mcp.WithString("end_date", mcp.Description("结束日期 YYYY-MM-DD")), ) } func (this *tool_get_memory_stats) Handl(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { uid, err := this.module.ResolveUID(ctx, request.GetString("uid", "")) if err != nil { return mcp.NewToolResultError(err.Error()), nil } start, end, errRange := resolveRange(request) if errRange != nil { return mcp.NewToolResultError(errRange.Error()), nil } type catCount struct { Category string State int32 N int64 } rows := make([]catCount, 0) if e := comm.ApplyMemoryQuery(mysql.Table(comm.TableMemoryItem), comm.MemoryQuery{Uid: uid, StartDate: start, EndDate: end}). Select("category, state, count(*) as n").Group("category, state"). Find(&rows).Error; e != nil { this.module.Errorln("get_memory_stats error:", e) return mcp.NewToolResultError("统计失败"), nil } var todoTotal, todoDone, alarmTotal, ideaCount int64 for _, r := range rows { switch r.Category { case comm.MemoryCatTodo: todoTotal += r.N if r.State == int32(pb.MemoryState_MemoryState_Done) { todoDone += r.N } case comm.MemoryCatAlarm: alarmTotal += r.N case comm.MemoryCatIdea: ideaCount += r.N } } // 花销**按币种分行**,不做汇率换算:把 CNY 和 JPY 加成一个数是错的。 type expRow struct { Currency string TotalCents int64 Count int64 } exps := make([]expRow, 0) if e := comm.ApplyMemoryQuery(mysql.Table(comm.TableMemoryItem), comm.MemoryQuery{Uid: uid, StartDate: start, EndDate: end, Categories: []string{comm.MemoryCatExpense}}). Select("currency, sum(amount_cents) as total_cents, count(*) as count"). Group("currency").Find(&exps).Error; e != nil { this.module.Errorln("get_memory_stats expense error:", e) return mcp.NewToolResultError("统计失败"), nil } expOut := make([]map[string]interface{}, 0, len(exps)) for _, e := range exps { expOut = append(expOut, map[string]interface{}{ "currency": e.Currency, // 同时给「分」和「元」:模型直接念「元」不容易错, // 但保留「分」以免它想自己做加减时用浮点数 "total_cents": e.TotalCents, "total": float64(e.TotalCents) / 100.0, "count": e.Count, }) } return mcp.NewToolResultText(utils.ToString(map[string]interface{}{ "result": "success", "start_date": start, "end_date": end, "todo_total": todoTotal, "todo_done": todoDone, "todo_undone": todoTotal - todoDone, "alarm_total": alarmTotal, "idea_count": ideaCount, "expenses": expOut, })), nil } // ---------- get_memory_report ---------- type tool_get_memory_report struct { cbase.ModuleCompBase module *Mcp } func (this *tool_get_memory_report) 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_get_memory_report) Start() (err error) { err = this.ModuleCompBase.Start() this.module.AddTool(ToolGroup_GLOBAL, this.Tool(), this.Handl) return } func (this *tool_get_memory_report) Tool() mcp.Tool { return mcp.NewTool("get_memory_report", mcp.WithDescription("获取用户的周报或月报(系统定期生成的总结)"), mcp.WithString("period_type", mcp.Description("week=周报 month=月报"), mcp.DefaultString("week"), ), mcp.WithString("period_key", mcp.Description("周期标识,周报形如 2026-W35,月报形如 2026-08;留空=最近一份"), ), ) } func (this *tool_get_memory_report) Handl(ctx context.Context, request mcp.CallToolRequest) (*mcp.CallToolResult, error) { uid, err := this.module.ResolveUID(ctx, request.GetString("uid", "")) if err != nil { return mcp.NewToolResultError(err.Error()), nil } ptype := strings.ToLower(strings.TrimSpace(request.GetString("period_type", comm.MemoryPeriodWeek))) if ptype != comm.MemoryPeriodWeek && ptype != comm.MemoryPeriodMonth { ptype = comm.MemoryPeriodWeek } pkey := strings.TrimSpace(request.GetString("period_key", "")) recs := make([]*pb.DBMemoryReport, 0) tx := mysql.Table(comm.TableMemoryReport). Where("uid = ? and period_type = ?", uid, ptype). // 只给已生成完的:还在生成中的报告 summary 是空的, // 给出去只会让模型编一段话填空。 Where("state = ?", int32(pb.MemoryReportState_MemoryReportState_Done)) if pkey != "" { tx = tx.Where("period_key = ?", pkey) } if e := tx.Order("period_end desc").Limit(1).Find(&recs).Error; e != nil { this.module.Errorln("get_memory_report error:", e) return mcp.NewToolResultError("查询失败"), nil } if len(recs) == 0 { return mcp.NewToolResultText(utils.ToString(map[string]interface{}{ "result": "empty", "message": "还没有生成过对应的报告", })), nil } r := recs[0] return mcp.NewToolResultText(utils.ToString(map[string]interface{}{ "result": "success", "period_type": r.PeriodType, "period_key": r.PeriodKey, "period_start": r.PeriodStart, "period_end": r.PeriodEnd, "summary": r.Summary, "stat": r.StatJson, })), nil } // ---------- 公共辅助 ---------- // resolveRange 解析时间范围:range 优先,其次 start_date/end_date,都没有则默认最近 7 天。 // // 提供 range 这种相对说法是刻意的:模型对「今天几号」的认知来自对话上下文, // 很不可靠,让它自己算 start_date 经常算错(SET_clock 把「下午 3:30」填成 // 03:30 是同一类问题)。 func resolveRange(request mcp.CallToolRequest) (start, end string, err error) { now := time.Now() if r := strings.TrimSpace(request.GetString("range", "")); r != "" { return comm.MemoryRelativeRange(r, now) } start = strings.TrimSpace(request.GetString("start_date", "")) end = strings.TrimSpace(request.GetString("end_date", "")) if start == "" && end == "" { return comm.FormatMemoryDate(now.AddDate(0, 0, -7)), comm.FormatMemoryDate(now.AddDate(0, 0, 7)), nil } if start != "" { if _, ok := comm.ParseMemoryDate(start); !ok { return "", "", errBadDate(start) } } if end != "" { if _, ok := comm.ParseMemoryDate(end); !ok { return "", "", errBadDate(end) } } return start, end, nil } func errBadDate(s string) error { return &badDateError{s} } type badDateError struct{ v string } func (e *badDateError) Error() string { return "日期格式应为 YYYY-MM-DD:" + e.v } // briefItems 只给模型必要字段。 // // 不把整行 DBMemoryItem 丢过去:那里面有 client_key、gen_round、extra 这类 // 内部字段,对回答问题毫无用处,只会占 token 并诱导模型去解释它们。 func briefItems(items []*pb.DBMemoryItem) []map[string]interface{} { out := make([]map[string]interface{}, 0, len(items)) for _, it := range items { m := map[string]interface{}{ "id": it.Id, "category": it.Category, "title": it.Title, "date": it.HappenDate, "done": it.State == pb.MemoryState_MemoryState_Done, } if it.HappenTime != "" { m["time"] = it.HappenTime } if it.Detail != "" { m["detail"] = it.Detail } if it.Owner != "" { m["owner"] = it.Owner } if it.DueRaw != "" { m["due"] = it.DueRaw } if it.Category == comm.MemoryCatExpense { m["amount"] = float64(it.AmountCents) / 100.0 m["currency"] = it.Currency } if !it.DateCertain { // 让模型知道这个日期是推断的,别把它当成用户确认过的安排来陈述 m["date_is_inferred"] = true } out = append(out, m) } return out }