Browse Source
## 先修鉴权:此前公网接口完全没有身份校验
authFromRequest 只把 Authorization 塞进 context,**全仓没有一处读回来,也没有
任何 JWT 校验**;三个用户工具(get_user_tasks / allhelp_task / cancel_user_task)
的 uid 都是 mcp.WithString("uid", Required()) —— 由大模型填进来的普通参数。
而 mcp 是独立进程、独立端口 7300,不经过 gateway 那套 parseToken + isInWhiteList
(那只管 /api/* 和 /web/*),又必须对百炼公网可达。所以这不是「内网接口没做鉴权」,
是公网接口完全没有鉴权:知道 uid 就能读改任意人的数据。
新增 modules/mcp/auth.go:
- 从 Authorization 解 JWT 取 uid,与 gateway/core.go 同一套口径(RegisteredClaims,
uid 在 ID 字段);兼容 "Bearer xxx" 前缀(gateway 收的是裸 token,标准 MCP
客户端会加前缀)。
- 参数里的 uid **只用于比对**,不一致直接拒。不静默改用会话 uid——调用方显然
误以为自己能指定用户,让它失败比让它以为成功了更安全。
- 解不出会话 uid 一律拒,**不回退到参数 uid**,那等于这层没做。
- TokenKey 未配置时拒绝所有用户数据类工具:配置漏了导致鉴权静默失效,
比工具不可用严重得多。TokenKey 必须与 gateway 同值,两边都读 GATEWAY_TOKEN_KEY。
存量三个工具同步改造,uid 参数保留但改为可选(兼容百炼后台已配好的工具定义)。
9 个测试守住这条线,其中三条是核心安全断言:参数 uid 与会话不符必须拒、
没有 token 必须拒而不是回退、未配 TokenKey 必须全拒。另有伪造签名、过期 token、
空 uid claim 的用例。
## 共享查询层:避免 home 与 mcp 各写一份
comm/memoryquery.go。mcp 从不 RpcCall、只裸查 MySQL,如果记忆项的过滤条件、
分类白名单、排序口径两边各写一遍,改了一边另一边不报错,只会让 EMAI 答的和
App 里显示的对不上——而且没人会发现。memory 模块的 listItems 也改走这里。
ApplyMemoryQuery 里 uid 为空时拼 "1 = 0" 而不是不加条件:少一个 uid 条件就是把
全库记忆项返回给调用方,而它有两个调用方,其中一个公网可达。
另提供 MemoryRelativeRange 把「本周/上月/最近7天」换算成日期区间给模型用——
模型对「今天几号」的认知来自对话上下文,很不可靠,让它自己算 start_date 经常
算错(SET_clock 把「下午 3:30」填成 03:30 是同一类问题)。
## 三个只读工具
get_memory_items / get_memory_stats / get_memory_report。
**只读是刻意的**:新增/修改/删除不放这里。写路径的校验、幂等(client_key)、
提醒时刻重算都在 home 的 memory 模块内,在 MCP 里复制一份必然漂移。
录入走端侧指令(tool_calls)、查询走 MCP,这个分工不变。
- 花销按币种分行返回,不做汇率换算;同时给「分」和「元」,模型念元不容易错,
留分是以免它想自己做加减时用浮点。
- get_memory_report 只返回已生成完的:还在生成中的 summary 是空的,
给出去只会让模型编一段话填空。
- 返回给模型的是精简字段,不是整行 DBMemoryItem——client_key/gen_round/extra
这些内部字段对回答问题没用,只占 token 还诱导模型去解释它们。
- date_certain=false 的项带上 date_is_inferred 标记,让模型别把推断出来的日期
当成用户确认过的安排来陈述。
- 模型填错分类时忽略该筛选当作「不限分类」,而不是整次查询失败——它多半会再填错一次。
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
main
12 changed files with 848 additions and 35 deletions
@ -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", |
|||
} |
|||
@ -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) |
|||
} |
|||
} |
|||
@ -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 |
|||
} |
|||
@ -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) |
|||
} |
|||
} |
|||
@ -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 |
|||
} |
|||
Loading…
Reference in new issue