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.
 
 
 
 
 
 

401 lines
13 KiB

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.ResolveUIDWithToken(ctx,
request.GetString("uid", ""), request.GetString("auth_token", ""))
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.ResolveUIDWithToken(ctx,
request.GetString("uid", ""), request.GetString("auth_token", ""))
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
}
}
// 重复闹钟补上周期内多出来的次数。口径与 home 的周期统计共用 comm 里那两个函数——
// 各写一份的话会变成「App 里的周报说 7 次、EMAI 说 1 次」,而且谁都不报错。
alarmTotal += int64(repeatAlarmExtras(uid, start, end))
// 花销**按币种分行**,不做汇率换算:把 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.ResolveUIDWithToken(ctx,
request.GetString("uid", ""), request.GetString("auth_token", ""))
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,
// ⚠️ 归一:happen_date 是 type:date 列 + parseTime=True,
// 读出来是 `2026-09-07T00:00:00+08:00`。原样交给模型,它回答时就会
// 把这一长串念出来,或者拿它跟用户说的「9 月 7 号」对不上。
"date": comm.NormalizeMemoryDate(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
}
// repeatAlarmExtras 重复闹钟在 [start,end] 里比库里那一行多出来的次数。
// 查不到或日期坏了一律当 0:统计少算一点,好过整个工具报错。
func repeatAlarmExtras(uid, start, end string) int {
s, ok1 := comm.ParseMemoryDate(start)
e, ok2 := comm.ParseMemoryDate(end)
if !ok1 || !ok2 {
return 0
}
items := make([]*pb.DBMemoryItem, 0)
if err := comm.ApplyRepeatCandidates(mysql.Table(comm.TableMemoryItem), uid, end,
[]string{comm.MemoryCatAlarm}).Find(&items).Error; err != nil {
safeLogErrorf("get_memory_stats 展开重复闹钟失败已忽略: %v", err)
return 0
}
total := 0
for _, it := range items {
anchor, ok := comm.ParseMemoryDate(it.HappenDate)
if !ok {
continue
}
total += comm.MemoryExtraOccurrences(it.RepeatRule, it.Weekday, anchor, s, e)
}
return total
}