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