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.
367 lines
12 KiB
367 lines
12 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.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
|
|
}
|
|
|