package handler import ( "errors" "testing" ) // parseRelay 是免配额通道的入口校验。白名单 + 强制幂等键这两条必须守住: // 前者防止 relay 变成任意字符串的后门,后者是「同一条上游消息只转一次」的基础。 func TestParseRelay(t *testing.T) { cases := []struct { name string kind, key string wantKind string wantKey string wantErr bool }{ {name: "都为空 = 普通自主发信,正常扣配额", kind: "", key: "", wantKind: "", wantKey: ""}, {name: "总结转发", kind: "summary", key: "msg_1", wantKind: "summary", wantKey: "msg_1"}, {name: "权限转发", kind: "permission", key: "per_1", wantKind: "permission", wantKey: "per_1"}, {name: "两端空白被裁掉", kind: " summary ", key: " msg_2 ", wantKind: "summary", wantKey: "msg_2"}, // 白名单外的类型必须拒:否则 relay:"anything" 就绕过了配额 {name: "未知类型", kind: "whatever", key: "k", wantErr: true}, // 没有幂等键就无法阻止同一条上游消息反复转发 {name: "缺幂等键", kind: "summary", key: "", wantErr: true}, {name: "只给了键没给类型", kind: "", key: "k", wantErr: true}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { kind, key, err := parseRelay(c.kind, c.key) if c.wantErr { if err == nil { t.Fatalf("期望报错,实际通过:kind=%q key=%q", kind, key) } return } if err != nil { t.Fatalf("意外报错: %v", err) } if kind != c.wantKind || key != c.wantKey { t.Fatalf("得到 (%q, %q),期望 (%q, %q)", kind, key, c.wantKind, c.wantKey) } }) } } func TestParseRelayRejectsOverlongKey(t *testing.T) { long := make([]byte, 161) for i := range long { long[i] = 'k' } if _, _, err := parseRelay("summary", string(long)); err == nil { t.Fatal("超长 relay_key 应被拒绝(列宽 160)") } } // 免配额类型是白名单,不是黑名单。新增一种转发时必须同时更新这里, // 免得悄悄多出一条不受审视的免费通道。 func TestRelayKindsIsExactlyTwo(t *testing.T) { want := map[string]bool{"permission": true, "summary": true} if len(relayKinds) != len(want) { t.Fatalf("免配额类型数量变了:%v。新增前请确认它确实是 harness 代劳而非模型自主发信", relayKinds) } for k := range want { if !relayKinds[k] { t.Fatalf("缺少免配额类型 %q", k) } } } // 报错必须是 400 而不是 500:这些都是调用方参数问题 func TestParseRelayErrorsAreBadRequest(t *testing.T) { for _, c := range [][2]string{{"whatever", "k"}, {"summary", ""}, {"", "k"}} { _, _, err := parseRelay(c[0], c[1]) if err == nil { t.Fatalf("(%q,%q) 应报错", c[0], c[1]) } var he httpError if !errors.As(err, &he) || he.status != 400 { t.Fatalf("(%q,%q) 的错误不是 400: %#v", c[0], c[1], err) } } }