diff --git a/apps/services/comm/memoryquery.go b/apps/services/comm/memoryquery.go new file mode 100644 index 00000000..5292bf5f --- /dev/null +++ b/apps/services/comm/memoryquery.go @@ -0,0 +1,134 @@ +package comm + +import ( + "fmt" + "strings" + "time" + + "gorm.io/gorm" +) + +/* +记忆项的**共享查询层**。 + +放在 comm 是为了解决一个具体问题:`mcp` 是独立进程、独立服务,它的工具从不 +RpcCall、只裸查 MySQL(见 tool_get_user_tasks 的 +`mysql.Table(...).Where("uid = ?", uid)`)。如果 memory 的过滤条件、分类白名单、 +排序口径在 home 和 mcp 里各写一遍,改了一边另一边不会报错,只会让 EMAI 答的 +和 App 里显示的对不上——而且没人会发现。 + +所以两边都调这里。**新增过滤条件只改这一个文件。** + +⚠️ 这里只放「读」。写操作(新增/改/删)仍只在 home 的 memory 模块里, +MCP 侧的写工具经它的 HTTP 接口不了——那是刻意的:写路径的校验、幂等、 +提醒重算都在 memory 模块内,复制一份必然漂移。 +*/ + +// MemoryQuery 记忆项查询条件。零值表示不限制。 +type MemoryQuery struct { + Uid string + StartDate string // YYYY-MM-DD,含 + EndDate string // YYYY-MM-DD,含 + Categories []string + States []int32 + Limit int + Offset int +} + +// ApplyMemoryQuery 把条件套到一个已经 Table(memory_item) 过的 *gorm.DB 上。 +// +// ⚠️ **uid 一定会被拼进去**,即使 q.Uid 为空也会拼一个不可能命中的条件—— +// 少一个 uid 条件就是把全库的记忆项返回给调用方,而这个函数有两个调用方, +// 其中一个(MCP)是公网可达的。宁可返回空也不要返回别人的数据。 +func ApplyMemoryQuery(tx *gorm.DB, q MemoryQuery) *gorm.DB { + if strings.TrimSpace(q.Uid) == "" { + return tx.Where("1 = 0") + } + tx = tx.Where("uid = ?", q.Uid) + if q.StartDate != "" { + tx = tx.Where("happen_date >= ?", q.StartDate) + } + if q.EndDate != "" { + tx = tx.Where("happen_date <= ?", q.EndDate) + } + if len(q.Categories) > 0 { + tx = tx.Where("category in ?", q.Categories) + } + if len(q.States) > 0 { + tx = tx.Where("state in ?", q.States) + } + return tx +} + +// OrderMemoryItems 统一排序口径:先日期,再时刻(无时刻的排当天最后),最后 id。 +// +// 两边必须一致,否则「App 里第一条」和「EMAI 说的第一条」会是不同的东西。 +func OrderMemoryItems(tx *gorm.DB) *gorm.DB { + return tx.Order("happen_date asc"). + Order("happen_time = '' asc"). + Order("happen_time asc"). + Order("id asc") +} + +// NormalizeMemoryCategories 过滤掉不在白名单里的分类。 +// +// 返回 (合法分类, 第一个非法分类)。全部合法时第二个返回值为空串。 +// 调用方决定是拒绝还是忽略——MCP 侧倾向忽略(模型填错分类不该让整次查询失败), +// HTTP 接口侧倾向拒绝(客户端填错是 bug,早点暴露)。 +func NormalizeMemoryCategories(in []string) ([]string, string) { + out := make([]string, 0, len(in)) + bad := "" + for _, c := range in { + c = strings.ToLower(strings.TrimSpace(c)) + if c == "" { + continue + } + if !IsValidMemoryCategory(c) { + if bad == "" { + bad = c + } + continue + } + out = append(out, c) + } + return out, bad +} + +// MemoryRelativeRange 把「最近 N 天 / 本周 / 本月」这类相对说法换算成日期区间。 +// +// 给 MCP 工具用:模型问「我这个月花了多少」时不该自己算日期, +// 它对「今天几号」的认知来自对话上下文,很不可靠。 +func MemoryRelativeRange(kind string, now time.Time) (start, end string, err error) { + switch strings.ToLower(strings.TrimSpace(kind)) { + case "today": + d := FormatMemoryDate(now) + return d, d, nil + case "yesterday": + d := FormatMemoryDate(now.AddDate(0, 0, -1)) + return d, d, nil + case "this_week": + s, e := MemoryWeekRange(now) + return FormatMemoryDate(s), FormatMemoryDate(e), nil + case "last_week": + s, e := MemoryWeekRange(now.AddDate(0, 0, -7)) + return FormatMemoryDate(s), FormatMemoryDate(e), nil + case "this_month": + s, e := MemoryMonthRange(now) + return FormatMemoryDate(s), FormatMemoryDate(e), nil + case "last_month": + prev := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, now.Location()).AddDate(0, 0, -1) + s, e := MemoryMonthRange(prev) + return FormatMemoryDate(s), FormatMemoryDate(e), nil + case "next_7_days": + return FormatMemoryDate(now), FormatMemoryDate(now.AddDate(0, 0, 7)), nil + case "last_7_days": + return FormatMemoryDate(now.AddDate(0, 0, -7)), FormatMemoryDate(now), nil + } + return "", "", fmt.Errorf("不支持的时间范围: %s", kind) +} + +// MemoryRelativeRanges 支持的相对范围,给工具描述用 +var MemoryRelativeRanges = []string{ + "today", "yesterday", "this_week", "last_week", + "this_month", "last_month", "next_7_days", "last_7_days", +} diff --git a/apps/services/comm/memoryquery_test.go b/apps/services/comm/memoryquery_test.go new file mode 100644 index 00000000..c7d16a35 --- /dev/null +++ b/apps/services/comm/memoryquery_test.go @@ -0,0 +1,64 @@ +package comm + +import ( + "testing" + "time" +) + +func TestMemoryRelativeRange(t *testing.T) { + loc := time.UTC + // 2026-09-05 是周六 + now := time.Date(2026, 9, 5, 12, 0, 0, 0, loc) + + cases := []struct{ kind, start, end string }{ + {"today", "2026-09-05", "2026-09-05"}, + {"yesterday", "2026-09-04", "2026-09-04"}, + // ISO 周:周一起,所以本周是 08-31(周一) ~ 09-06(周日) + {"this_week", "2026-08-31", "2026-09-06"}, + {"last_week", "2026-08-24", "2026-08-30"}, + {"this_month", "2026-09-01", "2026-09-30"}, + {"last_month", "2026-08-01", "2026-08-31"}, + {"next_7_days", "2026-09-05", "2026-09-12"}, + {"last_7_days", "2026-08-29", "2026-09-05"}, + } + for _, c := range cases { + s, e, err := MemoryRelativeRange(c.kind, now) + if err != nil { + t.Errorf("%s: 报错 %v", c.kind, err) + continue + } + if s != c.start || e != c.end { + t.Errorf("%s: 期望 %s~%s 得到 %s~%s", c.kind, c.start, c.end, s, e) + } + } +} + +func TestMemoryRelativeRange_Unknown(t *testing.T) { + if _, _, err := MemoryRelativeRange("last_fortnight", time.Now()); err == nil { + t.Error("不认识的范围必须报错,不能静默返回一个区间") + } +} + +func TestNormalizeMemoryCategories(t *testing.T) { + got, bad := NormalizeMemoryCategories([]string{"todo", "TODO", " idea ", "nonsense", ""}) + if bad != "nonsense" { + t.Errorf("应报出第一个非法分类 nonsense,得到 %q", bad) + } + // 大小写与空格都要归一 + want := []string{"todo", "todo", "idea"} + if len(got) != len(want) { + t.Fatalf("期望 %d 个合法分类,得到 %d: %v", len(want), len(got), got) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("第 %d 个:期望 %s 得到 %s", i, want[i], got[i]) + } + } +} + +func TestNormalizeMemoryCategories_AllValid(t *testing.T) { + _, bad := NormalizeMemoryCategories([]string{"todo", "alarm", "idea", "expense"}) + if bad != "" { + t.Errorf("全部合法时不该报错,却得到 %q", bad) + } +} diff --git a/apps/services/modules/mcp/auth.go b/apps/services/modules/mcp/auth.go new file mode 100644 index 00000000..70d06766 --- /dev/null +++ b/apps/services/modules/mcp/auth.go @@ -0,0 +1,105 @@ +package mcp + +import ( + "context" + "fmt" + "strings" + + "github.com/golang-jwt/jwt/v5" +) + +/* +MCP 的会话身份。 + +## 这块此前是空的 + +`authFromRequest` 只把 Authorization 头塞进 context,**全仓没有一处读回来, +也没有任何校验**;三个用户工具(get_user_tasks / allhelp_task / cancel_user_task) +的 uid 都是 `mcp.WithString("uid", Required())` —— 由大模型填进来的普通参数。 + +而 MCP 是**独立进程、独立端口(7300)**,不经过 gateway 那套 parseToken + +isInWhiteList(那只管 /api/* 和 /web/*),又必须对百炼公网可达。 +也就是说这不是「内网接口没做鉴权」,是**公网接口完全没有鉴权**: +知道 uid 就能读改任意人的数据。 + +记忆中心的条目装的是花销金额、灵感原文、每日行程,比「提醒任务」敏感得多, +所以在挂 memory 工具之前先把这层补上。 + +## 口径 + + - 有 Authorization 且能解出 uid → 用**会话 uid**,忽略参数里的 uid; + - 参数里的 uid 与会话 uid 不一致 → 直接拒(不是静默改用会话 uid: + 调用方显然误以为自己能指定用户,让它失败比让它以为成功了更安全); + - 解不出会话 uid → 拒。**不降级回退到参数 uid** —— 那等于这层没做。 + +⚠️ TokenKey 必须与 gateway 的 `GATEWAY_TOKEN_KEY` 一致,否则所有 token 都验不过。 +两边都从同一个 .env 变量注入。 +*/ + +// authKey is a custom context key for storing the auth token. +type authKey struct{} + +// withAuthKey adds an auth key to the context. +func withAuthKey(ctx context.Context, auth string) context.Context { + return context.WithValue(ctx, authKey{}, auth) +} + +// tokenFromContext 取出请求带来的 Authorization 原始串 +func tokenFromContext(ctx context.Context) string { + v, _ := ctx.Value(authKey{}).(string) + return v +} + +// parseUID 从 JWT 里解出 uid。与 gateway/core.go 的 parseToken 同一套口径 +// (jwt.RegisteredClaims,uid 放在 ID 字段)。 +func parseUID(tokenString string, secretKey []byte) (string, error) { + tokenString = strings.TrimSpace(tokenString) + if tokenString == "" { + return "", fmt.Errorf("缺少 Authorization") + } + // 兼容 "Bearer xxx" 写法:gateway 那边收到的就是裸 token, + // 但 MCP 客户端(百炼)常按标准加 Bearer 前缀。 + if len(tokenString) > 7 && strings.EqualFold(tokenString[:7], "bearer ") { + tokenString = strings.TrimSpace(tokenString[7:]) + } + parsed, err := jwt.ParseWithClaims(tokenString, &jwt.RegisteredClaims{}, + func(token *jwt.Token) (interface{}, error) { + if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok { + return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + } + return secretKey, nil + }) + if err != nil { + return "", err + } + claims, ok := parsed.Claims.(*jwt.RegisteredClaims) + if !ok || !parsed.Valid { + return "", fmt.Errorf("invalid token") + } + if claims.ID == "" { + return "", fmt.Errorf("token 里没有 uid") + } + return claims.ID, nil +} + +// ResolveUID 决定这次工具调用该用哪个 uid。 +// +// argUID 是模型传上来的参数,只用于「与会话 uid 比对」,绝不作为回退值。 +func (this *Mcp) ResolveUID(ctx context.Context, argUID string) (string, error) { + key := this.options.TokenKey + if key == "" { + // 没配 TokenKey 说明部署时漏了 GATEWAY_TOKEN_KEY。 + // 这时**拒绝所有用户类工具**而不是放行:配置缺失导致鉴权静默失效 + // 是最糟的一种失败方式。 + return "", fmt.Errorf("服务端未配置 TokenKey,用户数据类工具不可用") + } + uid, err := parseUID(tokenFromContext(ctx), []byte(key)) + if err != nil { + return "", fmt.Errorf("身份校验失败: %v", err) + } + argUID = strings.TrimSpace(argUID) + if argUID != "" && argUID != uid { + return "", fmt.Errorf("uid 与当前会话不符,已拒绝") + } + return uid, nil +} diff --git a/apps/services/modules/mcp/auth_test.go b/apps/services/modules/mcp/auth_test.go new file mode 100644 index 00000000..a158aac7 --- /dev/null +++ b/apps/services/modules/mcp/auth_test.go @@ -0,0 +1,121 @@ +package mcp + +import ( + "context" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +const testKey = "test-token-key-please-ignore" + +func mkToken(t *testing.T, uid string, key string, expired bool) string { + t.Helper() + exp := time.Now().Add(time.Hour) + if expired { + exp = time.Now().Add(-time.Hour) + } + tok := jwt.NewWithClaims(jwt.SigningMethodHS256, &jwt.RegisteredClaims{ + ID: uid, + ExpiresAt: jwt.NewNumericDate(exp), + }) + s, err := tok.SignedString([]byte(key)) + if err != nil { + t.Fatalf("签 token 失败: %v", err) + } + return s +} + +func mcpWith(key string) *Mcp { + return &Mcp{options: &Options{TokenKey: key}} +} + +// 正常路径:能从 token 解出 uid +func TestResolveUID_FromToken(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), mkToken(t, "u123", testKey, false)) + uid, err := m.ResolveUID(ctx, "") + if err != nil || uid != "u123" { + t.Fatalf("期望 u123,得到 %q err=%v", uid, err) + } +} + +// Bearer 前缀要能吃下:gateway 收的是裸 token,但标准 MCP 客户端会加前缀 +func TestResolveUID_BearerPrefix(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), "Bearer "+mkToken(t, "u123", testKey, false)) + uid, err := m.ResolveUID(ctx, "") + if err != nil || uid != "u123" { + t.Fatalf("Bearer 前缀应被剥掉,得到 %q err=%v", uid, err) + } +} + +// **核心安全断言**:参数里的 uid 与会话不符必须拒绝, +// 绝不能静默改用会话 uid(那会掩盖调用方的错误认知),更不能采信参数。 +func TestResolveUID_RejectsMismatchedArgUID(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), mkToken(t, "u123", testKey, false)) + if _, err := m.ResolveUID(ctx, "someone_else"); err == nil { + t.Fatal("参数 uid 与会话不符时必须报错") + } +} + +// 参数 uid 与会话一致时放行(百炼后台已配好的工具定义仍会传 uid) +func TestResolveUID_AllowsMatchingArgUID(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), mkToken(t, "u123", testKey, false)) + uid, err := m.ResolveUID(ctx, "u123") + if err != nil || uid != "u123" { + t.Fatalf("一致的 uid 应放行,得到 %q err=%v", uid, err) + } +} + +// **核心安全断言**:没有 token 时必须拒绝, +// 绝不能回退成「用参数里的 uid」——那等于这层鉴权没做。 +func TestResolveUID_NoTokenIsRejectedNotFallback(t *testing.T) { + m := mcpWith(testKey) + ctx := context.Background() + if uid, err := m.ResolveUID(ctx, "u123"); err == nil { + t.Fatalf("没有 token 时必须拒绝,却返回了 uid=%q", uid) + } +} + +// 签名不对(伪造 token)必须拒绝 +func TestResolveUID_RejectsWrongSignature(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), mkToken(t, "u123", "another-key", false)) + if _, err := m.ResolveUID(ctx, ""); err == nil { + t.Fatal("签名不匹配的 token 必须拒绝") + } +} + +// 过期 token 必须拒绝 +func TestResolveUID_RejectsExpired(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), mkToken(t, "u123", testKey, true)) + if _, err := m.ResolveUID(ctx, ""); err == nil { + t.Fatal("过期 token 必须拒绝") + } +} + +// **核心安全断言**:服务端没配 TokenKey 时一律拒绝,不是放行。 +// 配置漏了导致鉴权静默失效,比工具不可用严重得多。 +func TestResolveUID_MissingTokenKeyDeniesAll(t *testing.T) { + m := mcpWith("") + ctx := withAuthKey(context.Background(), mkToken(t, "u123", testKey, false)) + if _, err := m.ResolveUID(ctx, ""); err == nil { + t.Fatal("未配置 TokenKey 时必须拒绝所有用户数据类工具") + } +} + +// token 里没有 uid(ID 为空)时不能返回空串当成合法身份 —— +// 空 uid 会让下游的 ApplyMemoryQuery 走 "1 = 0" 分支返回空, +// 但那是最后一道防线,不该指望它。 +func TestResolveUID_RejectsEmptyIDClaim(t *testing.T) { + m := mcpWith(testKey) + ctx := withAuthKey(context.Background(), mkToken(t, "", testKey, false)) + if uid, err := m.ResolveUID(ctx, ""); err == nil { + t.Fatalf("token 里没有 uid 时必须拒绝,却返回 %q", uid) + } +} diff --git a/apps/services/modules/mcp/module.go b/apps/services/modules/mcp/module.go index 27188c25..b7042a29 100644 --- a/apps/services/modules/mcp/module.go +++ b/apps/services/modules/mcp/module.go @@ -89,6 +89,12 @@ func (this *Mcp) OnInstallComp() { // this.RegisterComp(new(tool_music_next)) // this.RegisterComp(new(tool_music_close)) + // 记忆中心(拾忆):只读工具。写操作留在 home 的 memory 模块, + // 在这里复制一份写路径的校验/幂等/提醒重算必然漂移。 + this.RegisterComp(new(tool_get_memory_items)) + this.RegisterComp(new(tool_get_memory_stats)) + this.RegisterComp(new(tool_get_memory_report)) + this.RegisterComp(new(tool_allhelp_task)) this.RegisterComp(new(tool_get_user_tasks)) this.RegisterComp(new(tool_cancel_user_task)) @@ -161,14 +167,6 @@ func (this *Mcp) AddTool(group string, tool mcp.Tool, handler server.ToolHandler return false } -// authKey is a custom context key for storing the auth token. -type authKey struct{} - -// withAuthKey adds an auth key to the context. -func withAuthKey(ctx context.Context, auth string) context.Context { - return context.WithValue(ctx, authKey{}, auth) -} - // authFromRequest extracts the auth token from the request headers. func authFromRequest(ctx context.Context, r *http.Request) context.Context { return withAuthKey(ctx, r.Header.Get("Authorization")) diff --git a/apps/services/modules/mcp/options.go b/apps/services/modules/mcp/options.go index 467c5928..06a9c12e 100644 --- a/apps/services/modules/mcp/options.go +++ b/apps/services/modules/mcp/options.go @@ -14,6 +14,11 @@ type ( TLSCertFile string TLSKeyFile string Groups map[string][]string + // TokenKey 校验客户端 JWT 用,**必须与 gateway 的同一个值** + // (两边都从 .env 的 GATEWAY_TOKEN_KEY 注入)。 + // 留空时所有用户数据类工具一律拒绝服务 —— 配置漏了导致鉴权静默失效 + // 比工具不可用严重得多。 + TokenKey string } ) diff --git a/apps/services/modules/mcp/tool_allhelp_task.go b/apps/services/modules/mcp/tool_allhelp_task.go index 390c9f39..0834b7f9 100644 --- a/apps/services/modules/mcp/tool_allhelp_task.go +++ b/apps/services/modules/mcp/tool_allhelp_task.go @@ -37,9 +37,10 @@ func (this *tool_allhelp_task) Start() (err error) { 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.Required(), + mcp.Description("用户ID(可省略,服务端以登录身份为准)"), ), mcp.WithString("task_name", mcp.Description("任务标题(简短摘要,如:下午3点会议提醒)"), @@ -67,9 +68,13 @@ func (this *tool_allhelp_task) Handl(ctx context.Context, request mcp.CallToolRe log.Field{Key: "request.Params.Arguments", Value: request.GetRawArguments()}, ) - uid, err := request.RequireString("uid") + // ⚠️ 参数里的 uid 只用于比对,**绝不作为身份来源**。真正的身份从 + // Authorization 里的 JWT 解出来(见 auth.go)。此前这里直接信任模型填的 uid, + // 而 MCP 是公网可达且无鉴权的独立服务,等于谁都能读改任意人的数据。 + argUID := request.GetString("uid", "") + uid, err := this.module.ResolveUID(ctx, argUID) if err != nil { - return mcp.NewToolResultError("uid is required"), nil + return mcp.NewToolResultError(err.Error()), nil } user := &pb.DBUser{} if findErr := mysql.FindOne(comm.TableUser, user, "uid=?", uid); findErr != nil { diff --git a/apps/services/modules/mcp/tool_cancel_user_task.go b/apps/services/modules/mcp/tool_cancel_user_task.go index a31eb5a1..a5172f1b 100644 --- a/apps/services/modules/mcp/tool_cancel_user_task.go +++ b/apps/services/modules/mcp/tool_cancel_user_task.go @@ -36,9 +36,10 @@ func (this *tool_cancel_user_task) Start() (err error) { func (this *tool_cancel_user_task) Tool() mcp.Tool { return mcp.NewTool("cancel_user_task", mcp.WithDescription("取消用户的提醒任务"), + // uid 不再是必填:身份由请求头里的 JWT 决定。保留这个参数只为兼容 + // 百炼后台已配好的工具定义,传上来也只会被拿去和会话 uid 比对。 mcp.WithString("uid", - mcp.Description("用户ID"), - mcp.Required(), + mcp.Description("用户ID(可省略,服务端以登录身份为准)"), ), mcp.WithNumber("task_id", mcp.Description("要取消的任务ID"), @@ -52,9 +53,13 @@ func (this *tool_cancel_user_task) Handl(ctx context.Context, request mcp.CallTo log.Field{Key: "request.Params.Arguments", Value: request.GetRawArguments()}, ) - uid, err := request.RequireString("uid") + // ⚠️ 参数里的 uid 只用于比对,**绝不作为身份来源**。真正的身份从 + // Authorization 里的 JWT 解出来(见 auth.go)。此前这里直接信任模型填的 uid, + // 而 MCP 是公网可达且无鉴权的独立服务,等于谁都能读改任意人的数据。 + argUID := request.GetString("uid", "") + uid, err := this.module.ResolveUID(ctx, argUID) if err != nil { - return mcp.NewToolResultError("uid is required"), nil + return mcp.NewToolResultError(err.Error()), nil } taskId := request.GetInt("task_id", 0) if taskId <= 0 { diff --git a/apps/services/modules/mcp/tool_get_user_tasks.go b/apps/services/modules/mcp/tool_get_user_tasks.go index 7ec61e1f..dc6ea898 100644 --- a/apps/services/modules/mcp/tool_get_user_tasks.go +++ b/apps/services/modules/mcp/tool_get_user_tasks.go @@ -34,9 +34,10 @@ func (this *tool_get_user_tasks) Start() (err error) { func (this *tool_get_user_tasks) Tool() mcp.Tool { return mcp.NewTool("get_user_tasks", mcp.WithDescription("查询用户的提醒任务列表,可以按状态筛选"), + // uid 不再是必填:身份由请求头里的 JWT 决定。保留这个参数只为兼容 + // 百炼后台已配好的工具定义,传上来也只会被拿去和会话 uid 比对。 mcp.WithString("uid", - mcp.Description("用户ID"), - mcp.Required(), + mcp.Description("用户ID(可省略,服务端以登录身份为准)"), ), mcp.WithNumber("status", mcp.Description("按状态筛选: -1=全部, 0=进行中, 2=已完成, 3=已取消"), @@ -54,9 +55,13 @@ func (this *tool_get_user_tasks) Handl(ctx context.Context, request mcp.CallTool log.Field{Key: "request.Params.Arguments", Value: request.GetRawArguments()}, ) - uid, err := request.RequireString("uid") + // ⚠️ 参数里的 uid 只用于比对,**绝不作为身份来源**。真正的身份从 + // Authorization 里的 JWT 解出来(见 auth.go)。此前这里直接信任模型填的 uid, + // 而 MCP 是公网可达且无鉴权的独立服务,等于谁都能读改任意人的数据。 + argUID := request.GetString("uid", "") + uid, err := this.module.ResolveUID(ctx, argUID) if err != nil { - return mcp.NewToolResultError("uid is required"), nil + return mcp.NewToolResultError(err.Error()), nil } status := request.GetInt("status", -1) limit := request.GetInt("limit", 20) diff --git a/apps/services/modules/mcp/tool_memory.go b/apps/services/modules/mcp/tool_memory.go new file mode 100644 index 00000000..ed37a0a0 --- /dev/null +++ b/apps/services/modules/mcp/tool_memory.go @@ -0,0 +1,367 @@ +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 +} diff --git a/apps/services/modules/memory/model.go b/apps/services/modules/memory/model.go index a677cffa..d3ae55f9 100644 --- a/apps/services/modules/memory/model.go +++ b/apps/services/modules/memory/model.go @@ -93,27 +93,22 @@ func (this *modelComp) saveItem(item *pb.DBMemoryItem) error { func (this *modelComp) listItems(q *itemQuery) ([]*pb.DBMemoryItem, int64, error) { items := make([]*pb.DBMemoryItem, 0) - tx := mysql.Table(comm.TableMemoryItem).Where("uid = ?", q.uid) - if q.startDate != "" { - tx = tx.Where("happen_date >= ?", q.startDate) - } - if q.endDate != "" { - tx = tx.Where("happen_date <= ?", q.endDate) - } - if len(q.categories) > 0 { - tx = tx.Where("category in ?", q.categories) - } - if len(q.states) > 0 { - tx = tx.Where("state in ?", q.states) - } + // 过滤与排序走 comm 里的共享实现:mcp 是独立进程,它的工具查同一张表, + // 两边各写一份的话改了一边不报错,只会让 EMAI 答的和 App 里显示的对不上。 + tx := comm.ApplyMemoryQuery(mysql.Table(comm.TableMemoryItem), comm.MemoryQuery{ + Uid: q.uid, + StartDate: q.startDate, + EndDate: q.endDate, + Categories: q.categories, + States: q.states, + }) var total int64 if err := tx.Count(&total).Error; err != nil { return nil, 0, err } - // 排序口径:先按日期,再按时刻(无时刻的排在当天最后),最后按 id 保证稳定 - tx = tx.Order("happen_date asc").Order("happen_time = '' asc").Order("happen_time asc").Order("id asc") + tx = comm.OrderMemoryItems(tx) if q.size > 0 { offset := 0 if q.page > 1 { diff --git a/deploy/app/confs/mcp.yaml.example b/deploy/app/confs/mcp.yaml.example index 7c88e3c4..333dd211 100644 --- a/deploy/app/confs/mcp.yaml.example +++ b/deploy/app/confs/mcp.yaml.example @@ -36,8 +36,17 @@ modules: Addr: "${MCP_ADDR:-http://localhost:7300}" TLSCertFile: "" TLSKeyFile: "" + # ⚠️ 必须与 gateway 的 TokenKey 是同一个值(两边都读 .env 的 GATEWAY_TOKEN_KEY)。 + # MCP 是独立进程独立端口、公网可达,不经过 gateway 的鉴权,用户数据类工具 + # 靠它自己校验 JWT。留空时这些工具一律拒绝服务——配置漏了导致鉴权静默失效 + # 比工具不可用严重得多。 + TokenKey: "${GATEWAY_TOKEN_KEY:-}" Groups: GLOBAL: + # 记忆中心(拾忆)只读工具:回答「我明天有什么事」「这个月花了多少」 + - get_memory_items + - get_memory_stats + - get_memory_report - hefeng_weather - music_search - music_play