feat: 实现 portal 提交契约与游标历史 (#25)
This commit is contained in:
@@ -2,8 +2,8 @@
|
||||
generated: true (请先修改 Gitea Wiki,禁止直接编辑本文件)
|
||||
wiki_page: Architecture-and-Code-Map
|
||||
wiki_url: https://git.ilapage.cn/OPC/chorus/wiki/Architecture-and-Code-Map.-
|
||||
wiki_revision: 8158e9e681f8c4a928a7bc5b744215bb4376cf05
|
||||
synchronized_at: 2026-08-22T03:42:41Z
|
||||
wiki_revision: abef747ecc01c1538c376d9c35a299d3bc4fb3c5
|
||||
synchronized_at: 2026-08-22T04:05:56Z
|
||||
<!-- gitea-wiki-mirror:end -->
|
||||
|
||||
# 架构与代码地图
|
||||
@@ -211,9 +211,9 @@ ClaimLease(只取得租约并记录 lease 事件)
|
||||
|
||||
#### portal 契约
|
||||
|
||||
- multipart 提交增加 `metadata` JSON,声明 `role_rule` 和每个文件 client_id/position/role/note;文件 part 必须与 metadata 一一对应。image_generate 不接受图片,image_edit 至少一张且恰好一个 primary;position 必须为从 0 开始的连续唯一值,role 仅为 primary/reference,note 与 role_rule 做 UTF-8/大小校验。
|
||||
- multipart 提交增加 `metadata` JSON,声明 `capability`、`role_rule` 和每个文件 client_id/position/role/note;每个文件 part 的字段名固定为 `files[<client_id>]`,必须与 `metadata.images[].client_id` 一一对应且各出现一次,不接受未声明文件。image_generate 不接受图片或 role_rule,image_edit 至少一张且恰好一个 primary;position 必须为从 0 开始的连续唯一值,role 仅为 primary/reference,note 与 role_rule 做 UTF-8/大小校验。
|
||||
- Prompt 覆盖顺序固定为 PromptTemplate 默认 role_rule → 非空用户 role_rule → 按 position 排序的图片 role/note;最终文本一次性写入 `rendered_prompt`,worker 不重新解释用户字段。
|
||||
- `GET /api/generations?limit=20&cursor=<opaque>` 使用带 HMAC 的 base64url `(created_at,id)` 游标,limit 有服务端上限;查询固定带 user_id,并按 `(created_at DESC,id DESC)` 取 `limit+1`。响应为 `items`、`next_cursor`、`has_more`;空/过期/篡改游标返回稳定 400,不退回 offset,也不泄露其他用户记录。
|
||||
- `GET /api/generations?limit=20&cursor=<opaque>` 使用带 HMAC 的 base64url `(created_at,id)` 游标,limit 有服务端上限;签名包含 user_id 作用域并使用独立域标签,密钥材料和有效期分别复用受保护的 session key 与 session TTL。查询固定带 user_id,并按 `(created_at DESC,id DESC)` 取 `limit+1`。响应为 `items`、`next_cursor`、`has_more`;空、过期、篡改或跨用户游标返回稳定 400,不退回 offset,也不泄露其他用户记录。
|
||||
|
||||
#### 必测矩阵
|
||||
|
||||
|
||||
@@ -2,8 +2,8 @@
|
||||
generated: true (请先修改 Gitea Wiki,禁止直接编辑本文件)
|
||||
wiki_page: Product-Requirements-Overview
|
||||
wiki_url: https://git.ilapage.cn/OPC/chorus/wiki/Product-Requirements-Overview.-
|
||||
wiki_revision: 71b97c2e46135b301934a43299e6346e4151df43
|
||||
synchronized_at: 2026-08-22T03:43:10Z
|
||||
wiki_revision: 5453c2433c2b5036f6ee7b6f86f1fdbe0092725c
|
||||
synchronized_at: 2026-08-22T04:06:22Z
|
||||
<!-- gitea-wiki-mirror:end -->
|
||||
|
||||
# 产品需求总览
|
||||
@@ -33,7 +33,7 @@ synchronized_at: 2026-08-22T03:43:10Z
|
||||
|
||||
## 当前需求索引
|
||||
|
||||
工单 [#1](https://git.ilapage.cn/OPC/chorus/issues/1) 已完成长期技术基线整理和验收。项目由 [Epic #3](https://git.ilapage.cn/OPC/chorus/issues/3) 统一跟踪;[MVP-0 #4](https://git.ilapage.cn/OPC/chorus/issues/4) 已于 2026-08-21 验收完成。[MVP-1 #19](https://git.ilapage.cn/OPC/chorus/issues/19) 的设计门禁已完成:技术设计 [#16](https://git.ilapage.cn/OPC/chorus/issues/16) 已于 2026-08-21 验收通过;管理端原型 [#17](https://git.ilapage.cn/OPC/chorus/issues/17) 的 `prototypes/17/v1/index.html` 已于 2026-08-21 通过用户验收,用户端原型 [#18](https://git.ilapage.cn/OPC/chorus/issues/18) 的 `prototypes/18/v1/index.html` 已于 2026-08-21 通过用户验收。生产实现与集成验收单元工单已建立:[#20](https://git.ilapage.cn/OPC/chorus/issues/20)~[#28](https://git.ilapage.cn/OPC/chorus/issues/28)。#20、#21、#22、#23 和 #28 已于 2026-08-21~2026-08-22 验收完成;#24 已于 2026-08-22 验收交付;#25~#27 按声明依赖待实施。其中 #23 已交付管理 API 基线;#24 按 #30 风险决策改为明文凭据与单条回显,并完成固定 go-admin-ui 页面、Casbin/菜单集成、脱敏审计和默认禁用的连通性探测门禁;portal 的 MVP-1 用户功能仍待 #25/#26 实现。
|
||||
工单 [#1](https://git.ilapage.cn/OPC/chorus/issues/1) 已完成长期技术基线整理和验收。项目由 [Epic #3](https://git.ilapage.cn/OPC/chorus/issues/3) 统一跟踪;[MVP-0 #4](https://git.ilapage.cn/OPC/chorus/issues/4) 已于 2026-08-21 验收完成。[MVP-1 #19](https://git.ilapage.cn/OPC/chorus/issues/19) 的设计门禁已完成:技术设计 [#16](https://git.ilapage.cn/OPC/chorus/issues/16) 已于 2026-08-21 验收通过;管理端原型 [#17](https://git.ilapage.cn/OPC/chorus/issues/17) 的 `prototypes/17/v1/index.html` 已于 2026-08-21 通过用户验收,用户端原型 [#18](https://git.ilapage.cn/OPC/chorus/issues/18) 的 `prototypes/18/v1/index.html` 已于 2026-08-21 通过用户验收。生产实现与集成验收单元工单已建立:[#20](https://git.ilapage.cn/OPC/chorus/issues/20)~[#28](https://git.ilapage.cn/OPC/chorus/issues/28)。#20、#21、#22、#23 和 #28 已于 2026-08-21~2026-08-22 验收完成;#24 已于 2026-08-22 验收交付;#25 已完成候选实现并等待用户验收,#26~#27 按声明依赖待实施。其中 #23 已交付管理 API 基线;#24 按 #30 风险决策改为明文凭据与单条回显,并完成固定 go-admin-ui 页面、Casbin/菜单集成、脱敏审计和默认禁用的连通性探测门禁;portal 的 MVP-1 提交契约与游标已由 #25 候选实现,用户页面仍待 #26 实现。
|
||||
|
||||
| 需求领域 | 用户与场景 | 需求状态 | MVP | 详细说明 | 实施工单 | 设计证据 |
|
||||
|---|---|---|---|---|---|---|
|
||||
@@ -42,7 +42,7 @@ synchronized_at: 2026-08-22T03:43:10Z
|
||||
| 默认 prompt template | 系统以数据配置而非硬编码合成提示词 | 已交付(#7/#13 已验收) | [MVP-0 #4](https://git.ilapage.cn/OPC/chorus/issues/4) | [业务规则](Business-Rules-and-Glossary.-) | [#7](https://git.ilapage.cn/OPC/chorus/issues/7) | 2026-08-20 已确认中性模板与 `{{.UserPrompt}}` |
|
||||
| 多 Provider 路由与故障转移 | 单上游故障时继续服务 | 已确认(#16 已验收,待实现) | [MVP-1 #19](https://git.ilapage.cn/OPC/chorus/issues/19) | [架构](Architecture-and-Code-Map.-)、[业务规则](Business-Rules-and-Glossary.-) | [#16 技术设计](https://git.ilapage.cn/OPC/chorus/issues/16) | #16 架构/API/数据/状态设计于 2026-08-21 已确认 |
|
||||
| 管理端配置与记录 | 运营配置模型、路由池并排障 | 已交付(#24 已验收) | [MVP-1 #19](https://git.ilapage.cn/OPC/chorus/issues/19) | 本页“管理端页面” | [#17 管理端原型](https://git.ilapage.cn/OPC/chorus/issues/17) | `prototypes/17/v1/index.html`,v1,2026-08-21 用户已确认 |
|
||||
| 图片角色编辑 | 用户编辑 role_rule、角色、备注与顺序 | 原型已确认(#18) | [MVP-1 #19](https://git.ilapage.cn/OPC/chorus/issues/19) | 业务规则“提示词与上传” | [#18 用户端原型](https://git.ilapage.cn/OPC/chorus/issues/18) | `prototypes/18/v1/index.html`,v1,2026-08-21 用户已确认 |
|
||||
| 图片角色编辑 | 用户编辑 role_rule、角色、备注与顺序 | 后端候选实现待验收(#25),页面待 #26 | [MVP-1 #19](https://git.ilapage.cn/OPC/chorus/issues/19) | 业务规则“提示词与上传” | [#18 用户端原型](https://git.ilapage.cn/OPC/chorus/issues/18) | `prototypes/18/v1/index.html`,v1,2026-08-21 用户已确认 |
|
||||
| API Key 与程序调用 | 外部程序提交与查询任务 | 已确认 | MVP-2 | 本页“程序化调用” | 待建 | API/权限设计 |
|
||||
| 限流、保留期、公开注册 | 治理滥用和磁盘;决定是否开放用户获取 | 待确认(阈值/流程) | MVP-2 或后续 | 业务规则“需要补充什么” | 待建 | 安全/运维/认证设计 |
|
||||
|
||||
@@ -85,7 +85,7 @@ MVP-0 非目标:
|
||||
|
||||
#### MVP-1:可用性与运营
|
||||
|
||||
汇总工单为 [#19](https://git.ilapage.cn/OPC/chorus/issues/19),设计门禁已完成,生产实现单元工单已建立;#24 已于 2026-08-22 验收交付:
|
||||
汇总工单为 [#19](https://git.ilapage.cn/OPC/chorus/issues/19),设计门禁已完成,生产实现单元工单已建立;#24 已于 2026-08-22 验收交付,#25 候选实现等待用户验收:
|
||||
|
||||
- [x] [#16](https://git.ilapage.cn/OPC/chorus/issues/16) 路由、Provider、迁移、API、状态和测试设计(2026-08-21 验收通过);
|
||||
- [x] [#17](https://git.ilapage.cn/OPC/chorus/issues/17) 固定 go-admin/go-admin-ui 的 CRUD 与三个定制页原型(`prototypes/17/v1/index.html`,2026-08-21 用户验收通过);
|
||||
@@ -269,3 +269,7 @@ MVP-2 的 API Key 存哈希、身份仍属于 `users`。提交、查询、幂等
|
||||
## #24 管理端交付状态(2026-08-22,已验收)
|
||||
|
||||
固定 go-admin-ui 已导入仓库并实现 Provider、ProviderModel、PromptTemplate、Users CRUD,以及路由池、Provider 健康/主动探测和生成详情页面。Provider Key 按 #30 明文落库并仅在单条页面主动回显;用户页按 #31 完全移除点数、余额和配额。实现已通过 Go build/vet/test、前端 lint/32 个单测/生产构建、MySQL 8 migration up/down/up、管理 API 集成测试和真实账号密码/菜单/API 联调;浏览器实例不可用,未生成生产页面自动化截图,用户于 2026-08-22 完成人工验收并确认 #24 通过。
|
||||
|
||||
## #25 portal 提交契约候选状态(2026-08-22,待验收)
|
||||
|
||||
portal 已按 #16/#18 实现活动路由与模板快照、image_generate/image_edit metadata、一一对应的图片角色备注、固定 Prompt 覆盖顺序,以及 user_id 作用域的 HMAC 游标历史。实现不调用上游,不包含用户点数,也不改变登录和页面交互;后者继续由 #26 负责。候选已通过 Go build/vet/test/race 和本机 MySQL 8.4.8 隔离集成测试,#25 保持待验收。
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/queue"
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
|
||||
passwordpkg "git.ilapage.cn/OPC/chorus/internal/platform/password"
|
||||
platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage"
|
||||
@@ -95,6 +96,10 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var databaseName string
|
||||
if err := db.Raw("SELECT DATABASE()").Scan(&databaseName).Error; err != nil || databaseName != "chorus_test" {
|
||||
t.Fatalf("portal integration requires chorus_test database, got %q: %v", databaseName, err)
|
||||
}
|
||||
sqlDB, _ := db.DB()
|
||||
defer sqlDB.Close()
|
||||
suffix := strconv.FormatInt(time.Now().UnixNano(), 10)
|
||||
@@ -109,14 +114,21 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
}
|
||||
templates := []model.PromptTemplate{
|
||||
{TemplateKey: "portal-text-" + suffix, Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Portal Text", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true},
|
||||
{TemplateKey: "portal-image-" + suffix, Kind: model.KindImage, APIType: model.APIImagesEdits, Capability: model.CapabilityImageEdit, Name: "Portal Image", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true},
|
||||
{TemplateKey: "portal-generate-" + suffix, Kind: model.KindImage, APIType: model.APIImages, Capability: model.CapabilityImageGenerate, Name: "Portal Generate", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true},
|
||||
{TemplateKey: "portal-edit-" + suffix, Kind: model.KindImage, APIType: model.APIImagesEdits, Capability: model.CapabilityImageEdit, Name: "Portal Edit", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "default role rule", Enabled: true},
|
||||
}
|
||||
if err := db.Create(&templates).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
routeIDs, providerID := createPortalRoutes(t, db, suffix, templates)
|
||||
defer func() {
|
||||
db.Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.Generation{})
|
||||
db.Where("id IN ?", []uint64{templates[0].ID, templates[1].ID}).Delete(&model.PromptTemplate{})
|
||||
db.Exec("DELETE FROM active_routes WHERE route_pool_id IN ?", routeIDs)
|
||||
db.Exec("DELETE FROM route_pools WHERE id IN ?", routeIDs)
|
||||
db.Exec("DELETE FROM provider_model_capabilities WHERE provider_model_id IN (SELECT id FROM provider_models WHERE provider_id=?)", providerID)
|
||||
db.Exec("DELETE FROM provider_models WHERE provider_id=?", providerID)
|
||||
db.Exec("DELETE FROM providers WHERE id=?", providerID)
|
||||
db.Where("id IN ?", []uint64{templates[0].ID, templates[1].ID, templates[2].ID}).Delete(&model.PromptTemplate{})
|
||||
db.Where("id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.User{})
|
||||
}()
|
||||
queueRepository, _ := queue.NewMySQLRepository(db)
|
||||
@@ -125,7 +137,8 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
generationService, err := service.New(db, queueRepository, local, service.Config{MaxPromptBytes: 1000, MaxImages: 3, MaxImageBytes: 200 << 10, MaxUploadBytes: 1 << 20, MaxImagePixels: 10000, HistoryLimit: 20})
|
||||
routeRepository, _ := corerouter.NewGORMRepository(db)
|
||||
generationService, err := service.New(db, queueRepository, local, routeRepository, service.Config{MaxPromptBytes: 1000, MaxImages: 3, MaxImageBytes: 200 << 10, MaxUploadBytes: 1 << 20, MaxImagePixels: 10000, HistoryLimit: 20, CursorKey: []byte(strings.Repeat("c", 32)), CursorTTL: time.Hour})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
@@ -188,7 +201,7 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
if err := db.First(&generation, textID).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if generation.Status != model.StatusPending || generation.ProviderModelID != nil || generation.AttemptCount != 0 {
|
||||
if generation.Status != model.StatusPending || generation.ProviderModelID != nil || generation.AttemptCount != 0 || generation.RoutePoolID == nil || generation.PromptTemplateID == nil || len(generation.RouteSnapshot) == 0 {
|
||||
t.Fatalf("synchronous submission touched worker fields: %#v", generation)
|
||||
}
|
||||
workspace := clientA.do(http.MethodGet, "/", nil, "")
|
||||
@@ -235,29 +248,87 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
t.Fatalf("succeeded text fragment=%d %s", succeededTextFragment.Code, succeededTextFragment.Body.String())
|
||||
}
|
||||
|
||||
generateKey := "generate-key-" + suffix
|
||||
generateResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, generateKey, "generate image", model.CapabilityImageGenerate, "", nil), lastMultipartType)
|
||||
if generateResponse.Code != http.StatusAccepted {
|
||||
t.Fatalf("image generate submit=%d %s", generateResponse.Code, generateResponse.Body.String())
|
||||
}
|
||||
generateID := responseID(t, generateResponse)
|
||||
var generated model.Generation
|
||||
if err := db.First(&generated, generateID).Error; err != nil || generated.RenderedPrompt != "generate image" || generated.RoleRule != nil {
|
||||
t.Fatalf("image generate=%#v err=%v", generated, err)
|
||||
}
|
||||
|
||||
imageKey := "image-key-" + suffix
|
||||
imageResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", []uploadPart{{"first.png", "image/png", testPNG(4, 4)}}), lastMultipartType)
|
||||
imageResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "user role rule", []uploadPart{{clientID: "primary-a", name: "first.png", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary, note: "keep edges"}}), lastMultipartType)
|
||||
if imageResponse.Code != 202 {
|
||||
t.Fatalf("image submit=%d %s", imageResponse.Code, imageResponse.Body.String())
|
||||
}
|
||||
imageID := responseID(t, imageResponse)
|
||||
var inputs []model.GenerationInput
|
||||
if err := db.Where("generation_id=?", imageID).Find(&inputs).Error; err != nil || len(inputs) != 1 || inputs[0].Role != model.RolePrimary {
|
||||
if err := db.Where("generation_id=?", imageID).Find(&inputs).Error; err != nil || len(inputs) != 1 || inputs[0].Role != model.RolePrimary || inputs[0].Note == nil || *inputs[0].Note != "keep edges" {
|
||||
t.Fatalf("inputs=%#v err=%v", inputs, err)
|
||||
}
|
||||
imageReplay := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", []uploadPart{{"second.png", "image/png", testPNG(4, 4)}}), lastMultipartType)
|
||||
var imageGeneration model.Generation
|
||||
if err := db.First(&imageGeneration, imageID).Error; err != nil || imageGeneration.RoleRule == nil || *imageGeneration.RoleRule != "user role rule" || !strings.Contains(imageGeneration.RenderedPrompt, "user role rule") || strings.Contains(imageGeneration.RenderedPrompt, "default role rule") || !strings.Contains(imageGeneration.RenderedPrompt, "position=0 role=primary note=keep edges") || imageGeneration.RoutePoolID == nil {
|
||||
t.Fatalf("image generation=%#v err=%v", imageGeneration, err)
|
||||
}
|
||||
if err := db.Table("route_pools").Where("id=?", routeIDs[2]).Update("enabled", false).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
imageReplay := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "changed rule", []uploadPart{{clientID: "primary-b", name: "second.png", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary}}), lastMultipartType)
|
||||
if imageReplay.Code != 200 || responseID(t, imageReplay) != imageID {
|
||||
t.Fatalf("image replay=%d %s", imageReplay.Code, imageReplay.Body.String())
|
||||
}
|
||||
if err := db.Table("route_pools").Where("id=?", routeIDs[2]).Update("enabled", true).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var inputCount int64
|
||||
db.Model(&model.GenerationInput{}).Where("generation_id=?", imageID).Count(&inputCount)
|
||||
if inputCount != 1 {
|
||||
t.Fatalf("replay wrote %d inputs", inputCount)
|
||||
}
|
||||
invalidResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, "invalid-key-"+suffix, "edit image", []uploadPart{{"fake.txt", "image/png", testPNG(4, 4)}}), lastMultipartType)
|
||||
invalidResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, "invalid-key-"+suffix, "edit image", model.CapabilityImageEdit, "", []uploadPart{{clientID: "bad-file", name: "fake.txt", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary}}), lastMultipartType)
|
||||
if invalidResponse.Code != 400 || !strings.Contains(invalidResponse.Body.String(), `"code":"invalid_image_extension"`) || !strings.Contains(invalidResponse.Body.String(), `"index":0`) || !strings.Contains(invalidResponse.Body.String(), `"name":"fake.txt"`) {
|
||||
t.Fatalf("invalid image response=%d %s", invalidResponse.Code, invalidResponse.Body.String())
|
||||
}
|
||||
historyResponse := clientA.do(http.MethodGet, "/api/generations?limit=2", nil, "")
|
||||
var firstHistory struct {
|
||||
Items []struct {
|
||||
ID uint64 `json:"id"`
|
||||
} `json:"items"`
|
||||
NextCursor string `json:"next_cursor"`
|
||||
HasMore bool `json:"has_more"`
|
||||
}
|
||||
if historyResponse.Code != 200 || json.Unmarshal(historyResponse.Body.Bytes(), &firstHistory) != nil || len(firstHistory.Items) != 2 || !firstHistory.HasMore || firstHistory.NextCursor == "" {
|
||||
t.Fatalf("first history=%d %s", historyResponse.Code, historyResponse.Body.String())
|
||||
}
|
||||
nextHistoryResponse := clientA.do(http.MethodGet, "/api/generations?limit=2&cursor="+firstHistory.NextCursor, nil, "")
|
||||
var nextHistory struct {
|
||||
Items []struct {
|
||||
ID uint64 `json:"id"`
|
||||
} `json:"items"`
|
||||
}
|
||||
if nextHistoryResponse.Code != 200 || json.Unmarshal(nextHistoryResponse.Body.Bytes(), &nextHistory) != nil || len(nextHistory.Items) == 0 || nextHistory.Items[0].ID == firstHistory.Items[0].ID || nextHistory.Items[0].ID == firstHistory.Items[1].ID {
|
||||
t.Fatalf("next history=%d %s", nextHistoryResponse.Code, nextHistoryResponse.Body.String())
|
||||
}
|
||||
tamperedBytes := []byte(firstHistory.NextCursor)
|
||||
tamperAt := len(tamperedBytes) / 2
|
||||
if tamperedBytes[tamperAt] == 'A' {
|
||||
tamperedBytes[tamperAt] = 'B'
|
||||
} else {
|
||||
tamperedBytes[tamperAt] = 'A'
|
||||
}
|
||||
tamperedCursor := string(tamperedBytes)
|
||||
if response := clientA.do(http.MethodGet, "/api/generations?limit=2&cursor="+tamperedCursor, nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
||||
t.Fatalf("tampered cursor=%d %s", response.Code, response.Body.String())
|
||||
}
|
||||
if response := clientA.do(http.MethodGet, "/api/generations?limit=21", nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_limit"`) {
|
||||
t.Fatalf("invalid limit=%d %s", response.Code, response.Body.String())
|
||||
}
|
||||
if response := clientA.do(http.MethodGet, "/api/generations?cursor=", nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
||||
t.Fatalf("empty cursor=%d %s", response.Code, response.Body.String())
|
||||
}
|
||||
|
||||
clientB := &apiClient{router: router}
|
||||
clientB.start(t)
|
||||
@@ -267,6 +338,9 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
if response := clientB.do(http.MethodGet, fmt.Sprintf("/api/generations/%d", imageID), nil, ""); response.Code != 404 {
|
||||
t.Fatalf("cross-user detail=%d", response.Code)
|
||||
}
|
||||
if response := clientB.do(http.MethodGet, "/api/generations?limit=2&cursor="+firstHistory.NextCursor, nil, ""); response.Code != 400 || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) {
|
||||
t.Fatalf("cross-user cursor=%d %s", response.Code, response.Body.String())
|
||||
}
|
||||
if response := clientB.do(http.MethodGet, fmt.Sprintf("/api/generations/%d/inputs/%d", imageID, inputs[0].ID), nil, ""); response.Code != 404 {
|
||||
t.Fatalf("cross-user input=%d", response.Code)
|
||||
}
|
||||
@@ -332,21 +406,29 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) {
|
||||
}
|
||||
|
||||
type uploadPart struct {
|
||||
name, mime string
|
||||
data []byte
|
||||
clientID, name, mime, note string
|
||||
data []byte
|
||||
position uint32
|
||||
role model.InputRole
|
||||
}
|
||||
|
||||
var lastMultipartType string
|
||||
|
||||
func multipartBody(t *testing.T, key, prompt string, parts []uploadPart) []byte {
|
||||
func multipartBody(t *testing.T, key, prompt string, capability model.Capability, roleRule string, parts []uploadPart) []byte {
|
||||
t.Helper()
|
||||
var body bytes.Buffer
|
||||
writer := multipart.NewWriter(&body)
|
||||
_ = writer.WriteField("idempotency_key", key)
|
||||
_ = writer.WriteField("prompt", prompt)
|
||||
metadata := service.ImageMetadata{Capability: capability, RoleRule: roleRule, Images: make([]service.ImageInput, 0, len(parts))}
|
||||
for _, item := range parts {
|
||||
metadata.Images = append(metadata.Images, service.ImageInput{ClientID: item.clientID, Position: item.position, Role: item.role, Note: item.note})
|
||||
}
|
||||
encodedMetadata, _ := json.Marshal(metadata)
|
||||
_ = writer.WriteField("metadata", string(encodedMetadata))
|
||||
for _, item := range parts {
|
||||
header := make(textproto.MIMEHeader)
|
||||
header.Set("Content-Disposition", fmt.Sprintf(`form-data; name="images"; filename="%s"`, item.name))
|
||||
header.Set("Content-Disposition", fmt.Sprintf(`form-data; name="files[%s]"; filename="%s"`, item.clientID, item.name))
|
||||
header.Set("Content-Type", item.mime)
|
||||
part, _ := writer.CreatePart(header)
|
||||
_, _ = part.Write(item.data)
|
||||
@@ -376,3 +458,57 @@ func testPNG(width, height int) []byte {
|
||||
_ = png.Encode(&body, img)
|
||||
return body.Bytes()
|
||||
}
|
||||
|
||||
func createPortalRoutes(t *testing.T, db *gorm.DB, suffix string, templates []model.PromptTemplate) ([]uint64, uint64) {
|
||||
t.Helper()
|
||||
tx := db.Begin()
|
||||
if tx.Error != nil {
|
||||
t.Fatal(tx.Error)
|
||||
}
|
||||
fail := func(err error) {
|
||||
t.Helper()
|
||||
tx.Rollback()
|
||||
t.Fatal(err)
|
||||
}
|
||||
slug := "portal-provider-" + suffix
|
||||
if err := tx.Exec("INSERT INTO providers (slug,name,base_url,auth_type,enabled) VALUES (?,?,?,'none',TRUE)", slug, "Portal Test", "https://example.invalid").Error; err != nil {
|
||||
fail(err)
|
||||
}
|
||||
var providerID uint64
|
||||
if err := tx.Raw("SELECT id FROM providers WHERE slug=?", slug).Scan(&providerID).Error; err != nil || providerID == 0 {
|
||||
fail(fmt.Errorf("load provider id: %w", err))
|
||||
}
|
||||
routeIDs := make([]uint64, 0, len(templates))
|
||||
for index, promptTemplate := range templates {
|
||||
modelIDValue := fmt.Sprintf("portal-model-%d-%s", index, suffix)
|
||||
if err := tx.Exec("INSERT INTO provider_models (provider_id,name,model_id,api_type,kind,extra_body,timeout_ms,weight,enabled) VALUES (?,?,?,?,?,JSON_OBJECT(),1000,100,TRUE)", providerID, "Portal Model", modelIDValue, promptTemplate.APIType, promptTemplate.Kind).Error; err != nil {
|
||||
fail(err)
|
||||
}
|
||||
var providerModelID uint64
|
||||
if err := tx.Raw("SELECT id FROM provider_models WHERE provider_id=? AND model_id=? AND api_type=?", providerID, modelIDValue, promptTemplate.APIType).Scan(&providerModelID).Error; err != nil || providerModelID == 0 {
|
||||
fail(fmt.Errorf("load provider model id: %w", err))
|
||||
}
|
||||
if err := tx.Exec("INSERT INTO provider_model_capabilities (provider_model_id,capability) VALUES (?,?)", providerModelID, promptTemplate.Capability).Error; err != nil {
|
||||
fail(err)
|
||||
}
|
||||
routeSlug := fmt.Sprintf("portal-route-%d-%s", index, suffix)
|
||||
if err := tx.Exec("INSERT INTO route_pools (slug,name,capability,prompt_template_id,max_failover,version,enabled) VALUES (?,?,?,?,0,1,TRUE)", routeSlug, "Portal Route", promptTemplate.Capability, promptTemplate.ID).Error; err != nil {
|
||||
fail(err)
|
||||
}
|
||||
var routeID uint64
|
||||
if err := tx.Raw("SELECT id FROM route_pools WHERE slug=?", routeSlug).Scan(&routeID).Error; err != nil || routeID == 0 {
|
||||
fail(fmt.Errorf("load route id: %w", err))
|
||||
}
|
||||
routeIDs = append(routeIDs, routeID)
|
||||
if err := tx.Exec("INSERT INTO route_pool_members (route_pool_id,provider_model_id,weight,failure_threshold,open_seconds,half_open_max,enabled,position) VALUES (?,?,100,3,60,1,TRUE,0)", routeID, providerModelID).Error; err != nil {
|
||||
fail(err)
|
||||
}
|
||||
if err := tx.Exec("INSERT INTO active_routes (capability,route_pool_id) VALUES (?,?)", promptTemplate.Capability, routeID).Error; err != nil {
|
||||
fail(err)
|
||||
}
|
||||
}
|
||||
if err := tx.Commit().Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return routeIDs, providerID
|
||||
}
|
||||
|
||||
+96
-15
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"html/template"
|
||||
@@ -14,6 +15,7 @@ import (
|
||||
"strings"
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
"git.ilapage.cn/OPC/chorus/portal/auth"
|
||||
"git.ilapage.cn/OPC/chorus/portal/service"
|
||||
"git.ilapage.cn/OPC/chorus/portal/session"
|
||||
@@ -177,23 +179,50 @@ func (h *Handler) submitImages(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
defer c.Request.MultipartForm.RemoveAll()
|
||||
headers := c.Request.MultipartForm.File["images"]
|
||||
uploads := make([]service.Upload, 0, len(headers))
|
||||
for index, header := range headers {
|
||||
var metadata service.ImageMetadata
|
||||
decoder := json.NewDecoder(strings.NewReader(c.PostForm("metadata")))
|
||||
decoder.DisallowUnknownFields()
|
||||
if err := decoder.Decode(&metadata); err != nil || decoder.Decode(&struct{}{}) != io.EOF {
|
||||
writeError(c, 400, "invalid_metadata", "image metadata is invalid")
|
||||
return
|
||||
}
|
||||
expectedFields := make(map[string]string, len(metadata.Images))
|
||||
for _, item := range metadata.Images {
|
||||
field := "files[" + item.ClientID + "]"
|
||||
if _, exists := expectedFields[field]; exists {
|
||||
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
|
||||
return
|
||||
}
|
||||
expectedFields[field] = item.ClientID
|
||||
}
|
||||
for field := range c.Request.MultipartForm.File {
|
||||
if _, ok := expectedFields[field]; !ok {
|
||||
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
|
||||
return
|
||||
}
|
||||
}
|
||||
uploads := make([]service.Upload, 0, len(expectedFields))
|
||||
for field, clientID := range expectedFields {
|
||||
headers := c.Request.MultipartForm.File[field]
|
||||
if len(headers) != 1 {
|
||||
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
|
||||
return
|
||||
}
|
||||
header := headers[0]
|
||||
file, err := header.Open()
|
||||
if err != nil {
|
||||
writeFileError(c, index, safeName(header.Filename), "invalid_image_content")
|
||||
writeFileError(c, int(metadataPosition(metadata, clientID)), safeName(header.Filename), "invalid_image_content")
|
||||
return
|
||||
}
|
||||
content, readErr := io.ReadAll(io.LimitReader(file, h.maxUploadBytes+1))
|
||||
file.Close()
|
||||
if readErr != nil {
|
||||
writeFileError(c, index, safeName(header.Filename), "invalid_image_content")
|
||||
writeFileError(c, int(metadataPosition(metadata, clientID)), safeName(header.Filename), "invalid_image_content")
|
||||
return
|
||||
}
|
||||
uploads = append(uploads, service.Upload{Name: header.Filename, DeclaredMIME: header.Header.Get("Content-Type"), Content: content})
|
||||
uploads = append(uploads, service.Upload{ClientID: clientID, Name: header.Filename, DeclaredMIME: header.Header.Get("Content-Type"), Content: content})
|
||||
}
|
||||
generation, created, err := h.service.SubmitImages(c.Request.Context(), currentSession(c).UserID, c.PostForm("idempotency_key"), c.PostForm("prompt"), uploads)
|
||||
generation, created, err := h.service.SubmitImages(c.Request.Context(), currentSession(c).UserID, c.PostForm("idempotency_key"), c.PostForm("prompt"), metadata, uploads)
|
||||
if err != nil {
|
||||
h.serviceError(c, err)
|
||||
return
|
||||
@@ -205,17 +234,40 @@ func (h *Handler) submitImages(c *gin.Context) {
|
||||
c.JSON(status, generationResponse(generation))
|
||||
}
|
||||
|
||||
func metadataPosition(metadata service.ImageMetadata, clientID string) uint32 {
|
||||
for _, item := range metadata.Images {
|
||||
if item.ClientID == clientID {
|
||||
return item.Position
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func (h *Handler) history(c *gin.Context) {
|
||||
rows, err := h.service.History(c.Request.Context(), currentSession(c).UserID)
|
||||
if err != nil {
|
||||
writeError(c, 500, "internal_error", "request could not be completed")
|
||||
limit := h.service.DefaultHistoryLimit()
|
||||
if raw, exists := c.GetQuery("limit"); exists {
|
||||
parsed, parseErr := strconv.Atoi(raw)
|
||||
if parseErr != nil {
|
||||
writeError(c, 400, "invalid_limit", "history limit is invalid")
|
||||
return
|
||||
}
|
||||
limit = parsed
|
||||
}
|
||||
cursor, hasCursor := c.GetQuery("cursor")
|
||||
if hasCursor && cursor == "" {
|
||||
writeError(c, 400, "invalid_cursor", "history cursor is invalid")
|
||||
return
|
||||
}
|
||||
items := make([]gin.H, 0, len(rows))
|
||||
for _, row := range rows {
|
||||
page, err := h.service.History(c.Request.Context(), currentSession(c).UserID, limit, cursor)
|
||||
if err != nil {
|
||||
h.serviceError(c, err)
|
||||
return
|
||||
}
|
||||
items := make([]gin.H, 0, len(page.Items))
|
||||
for _, row := range page.Items {
|
||||
items = append(items, generationResponse(row))
|
||||
}
|
||||
c.JSON(200, gin.H{"items": items})
|
||||
c.JSON(200, gin.H{"items": items, "next_cursor": page.NextCursor, "has_more": page.HasMore})
|
||||
}
|
||||
func (h *Handler) detail(c *gin.Context) {
|
||||
id, ok := uintParam(c, "id")
|
||||
@@ -241,9 +293,10 @@ func (h *Handler) detail(c *gin.Context) {
|
||||
}
|
||||
inputs := make([]gin.H, 0, len(detail.Inputs))
|
||||
for _, input := range detail.Inputs {
|
||||
inputs = append(inputs, gin.H{"id": input.ID, "position": input.Position, "role": input.Role, "name": input.OriginalName, "mime_type": input.MIMEType, "size_bytes": input.SizeBytes, "url": fmt.Sprintf("/api/generations/%d/inputs/%d", id, input.ID)})
|
||||
inputs = append(inputs, gin.H{"id": input.ID, "position": input.Position, "role": input.Role, "note": input.Note, "name": input.OriginalName, "mime_type": input.MIMEType, "size_bytes": input.SizeBytes, "url": fmt.Sprintf("/api/generations/%d/inputs/%d", id, input.ID)})
|
||||
}
|
||||
response := generationResponse(detail.Generation)
|
||||
response["role_rule"] = detail.Generation.RoleRule
|
||||
response["inputs"] = inputs
|
||||
response["outputs"] = outputs
|
||||
c.JSON(200, response)
|
||||
@@ -328,6 +381,34 @@ func (h *Handler) serviceError(c *gin.Context, err error) {
|
||||
writeError(c, 400, "invalid_idempotency_key", "idempotency key is invalid")
|
||||
case errors.Is(err, service.ErrImageCount):
|
||||
writeError(c, 400, "invalid_image_count", "image count is invalid")
|
||||
case errors.Is(err, service.ErrInvalidMetadata):
|
||||
writeError(c, 400, "invalid_metadata", "image metadata is invalid")
|
||||
case errors.Is(err, service.ErrInvalidCapability):
|
||||
writeError(c, 400, "invalid_capability", "generation capability is invalid")
|
||||
case errors.Is(err, service.ErrInvalidClientID):
|
||||
writeError(c, 400, "invalid_client_id", "image client id is invalid")
|
||||
case errors.Is(err, service.ErrImageMapping):
|
||||
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
|
||||
case errors.Is(err, service.ErrImagePosition):
|
||||
writeError(c, 400, "invalid_image_position", "image positions must be continuous")
|
||||
case errors.Is(err, service.ErrImageRole):
|
||||
writeError(c, 400, "invalid_image_role", "image role is invalid")
|
||||
case errors.Is(err, service.ErrImagePrimary):
|
||||
writeError(c, 400, "invalid_primary_count", "image edit requires exactly one primary image")
|
||||
case errors.Is(err, service.ErrInvalidRoleRule):
|
||||
writeError(c, 400, "invalid_role_rule", "image role rule is invalid")
|
||||
case errors.Is(err, service.ErrInvalidNote):
|
||||
writeError(c, 400, "invalid_image_note", "image note is invalid")
|
||||
case errors.Is(err, service.ErrInvalidLimit):
|
||||
writeError(c, 400, "invalid_limit", "history limit is invalid")
|
||||
case errors.Is(err, service.ErrInvalidCursor):
|
||||
writeError(c, 400, "invalid_cursor", "history cursor is invalid")
|
||||
case errors.Is(err, service.ErrCursorExpired):
|
||||
writeError(c, 400, "cursor_expired", "history cursor has expired")
|
||||
case errors.Is(err, corerouter.ErrRouteNotConfigured):
|
||||
writeError(c, 503, "route_not_configured", "generation route is not configured")
|
||||
case errors.Is(err, corerouter.ErrRouteUnavailable):
|
||||
writeError(c, 503, "route_unavailable", "generation route is unavailable")
|
||||
case errors.Is(err, service.ErrNotFound):
|
||||
writeError(c, 404, "not_found", "resource was not found")
|
||||
default:
|
||||
@@ -398,7 +479,7 @@ func (h *Handler) appPage(c *gin.Context) {
|
||||
c.Status(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
history, err := h.service.History(c.Request.Context(), state.UserID)
|
||||
history, err := h.service.RecentHistory(c.Request.Context(), state.UserID)
|
||||
if err != nil {
|
||||
c.Status(http.StatusInternalServerError)
|
||||
return
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
"git.ilapage.cn/OPC/chorus/portal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func TestMVP1ServiceErrorsHaveStableClientCodes(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
tests := []struct {
|
||||
err error
|
||||
code string
|
||||
want int
|
||||
}{
|
||||
{service.ErrInvalidMetadata, "invalid_metadata", 400},
|
||||
{service.ErrInvalidCapability, "invalid_capability", 400},
|
||||
{service.ErrInvalidClientID, "invalid_client_id", 400},
|
||||
{service.ErrImageMapping, "invalid_image_mapping", 400},
|
||||
{service.ErrImagePosition, "invalid_image_position", 400},
|
||||
{service.ErrImageRole, "invalid_image_role", 400},
|
||||
{service.ErrImagePrimary, "invalid_primary_count", 400},
|
||||
{service.ErrInvalidRoleRule, "invalid_role_rule", 400},
|
||||
{service.ErrInvalidNote, "invalid_image_note", 400},
|
||||
{service.ErrInvalidLimit, "invalid_limit", 400},
|
||||
{service.ErrInvalidCursor, "invalid_cursor", 400},
|
||||
{service.ErrCursorExpired, "cursor_expired", 400},
|
||||
{corerouter.ErrRouteNotConfigured, "route_not_configured", 503},
|
||||
{corerouter.ErrRouteUnavailable, "route_unavailable", 503},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.code, func(t *testing.T) {
|
||||
response := httptest.NewRecorder()
|
||||
context, _ := gin.CreateTestContext(response)
|
||||
(&Handler{}).serviceError(context, test.err)
|
||||
if response.Code != test.want || !strings.Contains(response.Body.String(), `"code":"`+test.code+`"`) {
|
||||
t.Fatalf("response=%d %s", response.Code, response.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
+6
-1
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/config"
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/queue"
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
safehttp "git.ilapage.cn/OPC/chorus/internal/platform/http"
|
||||
platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage"
|
||||
"git.ilapage.cn/OPC/chorus/portal/auth"
|
||||
@@ -62,6 +63,10 @@ func run() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
routeRepository, err := corerouter.NewGORMRepository(db)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
queueController := queue.NewController(queueRepository)
|
||||
httpClient, err := safehttp.New(safehttp.Config{Timeout: cfg.ProviderHTTPTimeout, MaxRedirects: 3})
|
||||
if err != nil {
|
||||
@@ -106,7 +111,7 @@ func run() error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
generationService, err := service.New(db, queueRepository, storage, service.Config{MaxPromptBytes: cfg.MaxPromptBytes, MaxImages: cfg.MaxImages, MaxImageBytes: cfg.MaxImageBytes, MaxUploadBytes: cfg.MaxUploadBytes, MaxImagePixels: cfg.MaxImagePixels, HistoryLimit: cfg.HistoryLimit})
|
||||
generationService, err := service.New(db, queueRepository, storage, routeRepository, service.Config{MaxPromptBytes: cfg.MaxPromptBytes, MaxImages: cfg.MaxImages, MaxImageBytes: cfg.MaxImageBytes, MaxUploadBytes: cfg.MaxUploadBytes, MaxImagePixels: cfg.MaxImagePixels, HistoryLimit: cfg.HistoryLimit, CursorKey: []byte(cfg.SessionKey), CursorTTL: cfg.SessionTTL})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrInvalidCursor = errors.New("invalid_cursor")
|
||||
ErrCursorExpired = errors.New("cursor_expired")
|
||||
)
|
||||
|
||||
const cursorPayloadBytes = 24
|
||||
|
||||
type historyCursor struct {
|
||||
CreatedAt time.Time
|
||||
ID uint64
|
||||
}
|
||||
|
||||
type cursorCodec struct {
|
||||
key []byte
|
||||
ttl time.Duration
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
func newCursorCodec(key []byte, ttl time.Duration) (cursorCodec, error) {
|
||||
if len(key) < 32 || ttl <= 0 {
|
||||
return cursorCodec{}, errors.New("cursor configuration is invalid")
|
||||
}
|
||||
return cursorCodec{key: append([]byte(nil), key...), ttl: ttl, now: time.Now}, nil
|
||||
}
|
||||
|
||||
func (c cursorCodec) encode(userID uint64, createdAt time.Time, id uint64) (string, error) {
|
||||
if userID == 0 || createdAt.IsZero() || id == 0 {
|
||||
return "", ErrInvalidCursor
|
||||
}
|
||||
payload := make([]byte, cursorPayloadBytes)
|
||||
binary.BigEndian.PutUint64(payload[0:8], uint64(createdAt.UTC().UnixMicro()))
|
||||
binary.BigEndian.PutUint64(payload[8:16], id)
|
||||
binary.BigEndian.PutUint64(payload[16:24], uint64(c.now().UTC().Add(c.ttl).UnixMicro()))
|
||||
signed := append(payload, c.sign(userID, payload)...)
|
||||
return base64.RawURLEncoding.EncodeToString(signed), nil
|
||||
}
|
||||
|
||||
func (c cursorCodec) decode(userID uint64, encoded string) (historyCursor, error) {
|
||||
if userID == 0 || encoded == "" {
|
||||
return historyCursor{}, ErrInvalidCursor
|
||||
}
|
||||
raw, err := base64.RawURLEncoding.DecodeString(encoded)
|
||||
if err != nil || len(raw) != cursorPayloadBytes+sha256.Size {
|
||||
return historyCursor{}, ErrInvalidCursor
|
||||
}
|
||||
payload, signature := raw[:cursorPayloadBytes], raw[cursorPayloadBytes:]
|
||||
if !hmac.Equal(signature, c.sign(userID, payload)) {
|
||||
return historyCursor{}, ErrInvalidCursor
|
||||
}
|
||||
expiresAt := int64(binary.BigEndian.Uint64(payload[16:24]))
|
||||
if !c.now().UTC().Before(time.UnixMicro(expiresAt).UTC()) {
|
||||
return historyCursor{}, ErrCursorExpired
|
||||
}
|
||||
createdMicros := int64(binary.BigEndian.Uint64(payload[0:8]))
|
||||
id := binary.BigEndian.Uint64(payload[8:16])
|
||||
if createdMicros <= 0 || id == 0 {
|
||||
return historyCursor{}, ErrInvalidCursor
|
||||
}
|
||||
return historyCursor{CreatedAt: time.UnixMicro(createdMicros).UTC(), ID: id}, nil
|
||||
}
|
||||
|
||||
func (c cursorCodec) sign(userID uint64, payload []byte) []byte {
|
||||
mac := hmac.New(sha256.New, c.key)
|
||||
_, _ = mac.Write([]byte("chorus:portal:history-cursor:v1\x00"))
|
||||
var scope [8]byte
|
||||
binary.BigEndian.PutUint64(scope[:], userID)
|
||||
_, _ = mac.Write(scope[:])
|
||||
_, _ = mac.Write(payload)
|
||||
return mac.Sum(nil)
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCursorRoundTripTamperScopeAndExpiry(t *testing.T) {
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
codec, err := newCursorCodec([]byte(strings.Repeat("k", 32)), time.Hour)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
codec.now = func() time.Time { return now }
|
||||
createdAt := now.Add(-time.Minute).Truncate(time.Microsecond)
|
||||
encoded, err := codec.encode(7, createdAt, 42)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
decoded, err := codec.decode(7, encoded)
|
||||
if err != nil || !decoded.CreatedAt.Equal(createdAt) || decoded.ID != 42 {
|
||||
t.Fatalf("decode = %#v, %v", decoded, err)
|
||||
}
|
||||
if _, err := codec.decode(8, encoded); !errors.Is(err, ErrInvalidCursor) {
|
||||
t.Fatalf("cross-user error = %v", err)
|
||||
}
|
||||
raw, _ := base64.RawURLEncoding.DecodeString(encoded)
|
||||
raw[len(raw)-1] ^= 1
|
||||
tampered := base64.RawURLEncoding.EncodeToString(raw)
|
||||
if _, err := codec.decode(7, tampered); !errors.Is(err, ErrInvalidCursor) {
|
||||
t.Fatalf("tamper error = %v", err)
|
||||
}
|
||||
codec.now = func() time.Time { return now.Add(time.Hour) }
|
||||
if _, err := codec.decode(7, encoded); !errors.Is(err, ErrCursorExpired) {
|
||||
t.Fatalf("expiry error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCursorRejectsInvalidValues(t *testing.T) {
|
||||
if _, err := newCursorCodec([]byte("short"), time.Hour); err == nil {
|
||||
t.Fatal("short key accepted")
|
||||
}
|
||||
codec, _ := newCursorCodec([]byte(strings.Repeat("k", 32)), time.Hour)
|
||||
if _, err := codec.encode(0, time.Now(), 1); !errors.Is(err, ErrInvalidCursor) {
|
||||
t.Fatalf("zero user error = %v", err)
|
||||
}
|
||||
if _, err := codec.decode(1, ""); !errors.Is(err, ErrInvalidCursor) {
|
||||
t.Fatalf("empty error = %v", err)
|
||||
}
|
||||
}
|
||||
+121
-23
@@ -10,11 +10,13 @@ import (
|
||||
"net/http"
|
||||
"path"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/generate"
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/queue"
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage"
|
||||
_ "golang.org/x/image/webp"
|
||||
"gorm.io/gorm"
|
||||
@@ -30,6 +32,16 @@ var (
|
||||
ErrImageExtension = errors.New("invalid_image_extension")
|
||||
ErrImageDecode = errors.New("invalid_image_content")
|
||||
ErrImagePixels = errors.New("image_pixels_exceeded")
|
||||
ErrInvalidMetadata = errors.New("invalid_metadata")
|
||||
ErrInvalidCapability = errors.New("invalid_capability")
|
||||
ErrInvalidClientID = errors.New("invalid_client_id")
|
||||
ErrImageMapping = errors.New("invalid_image_mapping")
|
||||
ErrImagePosition = errors.New("invalid_image_position")
|
||||
ErrImageRole = errors.New("invalid_image_role")
|
||||
ErrImagePrimary = errors.New("invalid_primary_count")
|
||||
ErrInvalidRoleRule = errors.New("invalid_role_rule")
|
||||
ErrInvalidNote = errors.New("invalid_image_note")
|
||||
ErrInvalidLimit = errors.New("invalid_limit")
|
||||
ErrNotFound = errors.New("not_found")
|
||||
)
|
||||
|
||||
@@ -41,16 +53,20 @@ type Config struct {
|
||||
MaxImageBytes, MaxUploadBytes int64
|
||||
MaxImagePixels uint64
|
||||
HistoryLimit int
|
||||
CursorKey []byte
|
||||
CursorTTL time.Duration
|
||||
}
|
||||
type Service struct {
|
||||
db *gorm.DB
|
||||
creator PreparedCreator
|
||||
storage corestorage.Store
|
||||
routes corerouter.SnapshotRepository
|
||||
cursors cursorCodec
|
||||
config Config
|
||||
}
|
||||
type Upload struct {
|
||||
Name, DeclaredMIME string
|
||||
Content []byte
|
||||
ClientID, Name, DeclaredMIME string
|
||||
Content []byte
|
||||
}
|
||||
type FileError struct {
|
||||
Index int
|
||||
@@ -61,55 +77,87 @@ type FileError struct {
|
||||
func (e *FileError) Error() string { return e.Code.Error() }
|
||||
func (e *FileError) Unwrap() error { return e.Code }
|
||||
|
||||
func New(db *gorm.DB, creator PreparedCreator, storage corestorage.Store, config Config) (*Service, error) {
|
||||
func New(db *gorm.DB, creator PreparedCreator, storage corestorage.Store, routes corerouter.SnapshotRepository, config Config) (*Service, error) {
|
||||
if db == nil || creator == nil || storage == nil || config.MaxPromptBytes <= 0 || config.MaxImages <= 0 || config.MaxImageBytes <= 0 || config.MaxUploadBytes < config.MaxImageBytes || config.MaxImagePixels == 0 || config.HistoryLimit <= 0 {
|
||||
return nil, errors.New("portal service configuration is invalid")
|
||||
}
|
||||
return &Service{db: db, creator: creator, storage: storage, config: config}, nil
|
||||
if routes == nil {
|
||||
return nil, errors.New("portal route repository is required")
|
||||
}
|
||||
cursors, err := newCursorCodec(config.CursorKey, config.CursorTTL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Service{db: db, creator: creator, storage: storage, routes: routes, cursors: cursors, config: config}, nil
|
||||
}
|
||||
|
||||
func (s *Service) SubmitText(ctx context.Context, userID uint64, idempotencyKey, prompt string) (model.Generation, bool, error) {
|
||||
if err := s.validateCommon(userID, idempotencyKey, prompt); err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
rendered, err := s.render(ctx, model.KindText, prompt)
|
||||
if existing, found, err := s.existingGeneration(ctx, userID, idempotencyKey); err != nil || found {
|
||||
return existing, false, err
|
||||
}
|
||||
snapshot, promptTemplate, err := s.submissionRoute(ctx, model.CapabilityText)
|
||||
if err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
rendered, err := generate.RenderPrompt(promptTemplate.TemplateText, prompt)
|
||||
if err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
generation := model.Generation{UserID: userID, Kind: model.KindText, IdempotencyKey: idempotencyKey, UserPrompt: prompt, RenderedPrompt: rendered}
|
||||
if err := corerouter.ApplySnapshot(&generation, snapshot); err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
created, err := s.creator.CreateIdempotentPrepared(ctx, &generation, func(uint64) ([]model.GenerationInput, error) { return nil, nil })
|
||||
return generation, created, err
|
||||
}
|
||||
|
||||
func (s *Service) SubmitImages(ctx context.Context, userID uint64, idempotencyKey, prompt string, uploads []Upload) (model.Generation, bool, error) {
|
||||
func (s *Service) SubmitImages(ctx context.Context, userID uint64, idempotencyKey, prompt string, metadata ImageMetadata, uploads []Upload) (model.Generation, bool, error) {
|
||||
if err := s.validateCommon(userID, idempotencyKey, prompt); err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
validated, err := s.validateImages(uploads)
|
||||
if existing, found, err := s.existingGeneration(ctx, userID, idempotencyKey); err != nil || found {
|
||||
return existing, false, err
|
||||
}
|
||||
validated, roleRule, err := s.validateImageSubmission(metadata, uploads)
|
||||
if err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
rendered, err := s.render(ctx, model.KindImage, prompt)
|
||||
snapshot, promptTemplate, err := s.submissionRoute(ctx, metadata.Capability)
|
||||
if err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
base, err := generate.RenderPrompt(promptTemplate.TemplateText, prompt)
|
||||
if err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
rendered := composeImagePrompt(base, promptTemplate.DefaultRoleRule, roleRule, validated)
|
||||
generation := model.Generation{UserID: userID, Kind: model.KindImage, IdempotencyKey: idempotencyKey, UserPrompt: prompt, RenderedPrompt: rendered}
|
||||
if roleRule != "" {
|
||||
generation.RoleRule = &roleRule
|
||||
}
|
||||
if err := corerouter.ApplySnapshot(&generation, snapshot); err != nil {
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
var saved []string
|
||||
created, err := s.creator.CreateIdempotentPrepared(ctx, &generation, func(generationID uint64) ([]model.GenerationInput, error) {
|
||||
rows := make([]model.GenerationInput, 0, len(validated))
|
||||
for index, item := range validated {
|
||||
key := fmt.Sprintf("inputs/%d/%d/%d", userID, generationID, index+1)
|
||||
for _, item := range validated {
|
||||
key := fmt.Sprintf("inputs/%d/%d/%d", userID, generationID, item.Position+1)
|
||||
object, putErr := s.storage.Put(ctx, corestorage.PutRequest{Key: key, OwnerID: userID, GenerationID: generationID, ContentType: item.MIME, Source: bytes.NewReader(item.Content)})
|
||||
if putErr != nil {
|
||||
return nil, putErr
|
||||
}
|
||||
saved = append(saved, key)
|
||||
width, height := uint32(item.Width), uint32(item.Height)
|
||||
role := model.RoleReference
|
||||
if index == 0 {
|
||||
role = model.RolePrimary
|
||||
var note *string
|
||||
if item.Note != "" {
|
||||
noteValue := item.Note
|
||||
note = ¬eValue
|
||||
}
|
||||
rows = append(rows, model.GenerationInput{Position: uint32(index), Role: role, OriginalName: item.Name, MIMEType: item.MIME, StorageKey: key, SizeBytes: uint64(object.Size), Width: &width, Height: &height})
|
||||
rows = append(rows, model.GenerationInput{Position: item.Position, Role: item.Role, Note: note, OriginalName: item.Name, MIMEType: item.MIME, StorageKey: key, SizeBytes: uint64(object.Size), Width: &width, Height: &height})
|
||||
}
|
||||
return rows, nil
|
||||
})
|
||||
@@ -125,6 +173,9 @@ type validatedImage struct {
|
||||
Name, MIME string
|
||||
Content []byte
|
||||
Width, Height int
|
||||
Position uint32
|
||||
Role model.InputRole
|
||||
Note string
|
||||
}
|
||||
|
||||
func (s *Service) validateImages(uploads []Upload) ([]validatedImage, error) {
|
||||
@@ -184,13 +235,17 @@ func (s *Service) validateCommon(userID uint64, key, prompt string) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (s *Service) render(ctx context.Context, kind model.GenerationKind, prompt string) (string, error) {
|
||||
var template model.PromptTemplate
|
||||
err := s.db.WithContext(ctx).Where("kind=? AND enabled=TRUE", kind).Order("version DESC, id DESC").First(&template).Error
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("load prompt template: %w", err)
|
||||
|
||||
func (s *Service) existingGeneration(ctx context.Context, userID uint64, key string) (model.Generation, bool, error) {
|
||||
var generation model.Generation
|
||||
err := s.db.WithContext(ctx).Where("user_id=? AND idempotency_key=?", userID, key).First(&generation).Error
|
||||
if err == nil {
|
||||
return generation, true, nil
|
||||
}
|
||||
return generate.RenderPrompt(template.TemplateText, prompt)
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return model.Generation{}, false, nil
|
||||
}
|
||||
return model.Generation{}, false, err
|
||||
}
|
||||
|
||||
type Detail struct {
|
||||
@@ -217,10 +272,53 @@ func (s *Service) ByID(ctx context.Context, userID, generationID uint64) (Detail
|
||||
}
|
||||
return Detail{Generation: generation, Inputs: inputs, Outputs: outputs}, nil
|
||||
}
|
||||
func (s *Service) History(ctx context.Context, userID uint64) ([]model.Generation, error) {
|
||||
|
||||
type HistoryPage struct {
|
||||
Items []model.Generation
|
||||
NextCursor string
|
||||
HasMore bool
|
||||
}
|
||||
|
||||
func (s *Service) History(ctx context.Context, userID uint64, limit int, cursor string) (HistoryPage, error) {
|
||||
if userID == 0 || limit <= 0 || limit > s.config.HistoryLimit {
|
||||
return HistoryPage{}, ErrInvalidLimit
|
||||
}
|
||||
query := s.db.WithContext(ctx).Where("user_id=?", userID)
|
||||
if cursor != "" {
|
||||
decoded, err := s.cursors.decode(userID, cursor)
|
||||
if err != nil {
|
||||
return HistoryPage{}, err
|
||||
}
|
||||
query = query.Where("(created_at < ? OR (created_at = ? AND id < ?))", decoded.CreatedAt, decoded.CreatedAt, decoded.ID)
|
||||
}
|
||||
var rows []model.Generation
|
||||
err := s.db.WithContext(ctx).Where("user_id=?", userID).Order("created_at DESC,id DESC").Limit(s.config.HistoryLimit).Find(&rows).Error
|
||||
return rows, err
|
||||
err := query.Order("created_at DESC,id DESC").Limit(limit + 1).Find(&rows).Error
|
||||
if err != nil {
|
||||
return HistoryPage{}, err
|
||||
}
|
||||
page := HistoryPage{Items: rows}
|
||||
if len(rows) > limit {
|
||||
page.HasMore = true
|
||||
page.Items = rows[:limit]
|
||||
last := page.Items[len(page.Items)-1]
|
||||
page.NextCursor, err = s.cursors.encode(userID, last.CreatedAt, last.ID)
|
||||
if err != nil {
|
||||
return HistoryPage{}, err
|
||||
}
|
||||
}
|
||||
return page, nil
|
||||
}
|
||||
|
||||
func (s *Service) RecentHistory(ctx context.Context, userID uint64) ([]model.Generation, error) {
|
||||
page, err := s.History(ctx, userID, s.config.HistoryLimit, "")
|
||||
return page.Items, err
|
||||
}
|
||||
|
||||
func (s *Service) DefaultHistoryLimit() int {
|
||||
if s.config.HistoryLimit < 20 {
|
||||
return s.config.HistoryLimit
|
||||
}
|
||||
return 20
|
||||
}
|
||||
|
||||
func (s *Service) User(ctx context.Context, userID uint64) (model.User, error) {
|
||||
|
||||
@@ -0,0 +1,149 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
var clientIDPattern = regexp.MustCompile(`^[A-Za-z0-9_-]{1,64}$`)
|
||||
|
||||
type ImageMetadata struct {
|
||||
Capability model.Capability `json:"capability"`
|
||||
RoleRule string `json:"role_rule"`
|
||||
Images []ImageInput `json:"images"`
|
||||
}
|
||||
|
||||
type ImageInput struct {
|
||||
ClientID string `json:"client_id"`
|
||||
Position uint32 `json:"position"`
|
||||
Role model.InputRole `json:"role"`
|
||||
Note string `json:"note"`
|
||||
}
|
||||
|
||||
func (s *Service) validateImageSubmission(metadata ImageMetadata, uploads []Upload) ([]validatedImage, string, error) {
|
||||
if metadata.Capability != model.CapabilityImageGenerate && metadata.Capability != model.CapabilityImageEdit {
|
||||
return nil, "", ErrInvalidCapability
|
||||
}
|
||||
roleRule := strings.TrimSpace(metadata.RoleRule)
|
||||
if !utf8.ValidString(metadata.RoleRule) || len(metadata.RoleRule) > s.config.MaxPromptBytes {
|
||||
return nil, "", ErrInvalidRoleRule
|
||||
}
|
||||
if metadata.Capability == model.CapabilityImageGenerate {
|
||||
if len(metadata.Images) != 0 || len(uploads) != 0 || roleRule != "" {
|
||||
return nil, "", ErrImageCount
|
||||
}
|
||||
return nil, "", nil
|
||||
}
|
||||
if len(metadata.Images) == 0 || len(metadata.Images) > s.config.MaxImages || len(uploads) != len(metadata.Images) {
|
||||
return nil, "", ErrImageCount
|
||||
}
|
||||
|
||||
uploadByID := make(map[string]Upload, len(uploads))
|
||||
for _, upload := range uploads {
|
||||
if !clientIDPattern.MatchString(upload.ClientID) {
|
||||
return nil, "", ErrInvalidClientID
|
||||
}
|
||||
if _, exists := uploadByID[upload.ClientID]; exists {
|
||||
return nil, "", ErrImageMapping
|
||||
}
|
||||
uploadByID[upload.ClientID] = upload
|
||||
}
|
||||
|
||||
inputs := append([]ImageInput(nil), metadata.Images...)
|
||||
sort.Slice(inputs, func(i, j int) bool { return inputs[i].Position < inputs[j].Position })
|
||||
seenIDs := make(map[string]struct{}, len(inputs))
|
||||
primaryCount := 0
|
||||
orderedUploads := make([]Upload, 0, len(inputs))
|
||||
for index, input := range inputs {
|
||||
if !clientIDPattern.MatchString(input.ClientID) {
|
||||
return nil, "", ErrInvalidClientID
|
||||
}
|
||||
if _, exists := seenIDs[input.ClientID]; exists {
|
||||
return nil, "", ErrImageMapping
|
||||
}
|
||||
seenIDs[input.ClientID] = struct{}{}
|
||||
if input.Position != uint32(index) {
|
||||
return nil, "", ErrImagePosition
|
||||
}
|
||||
if input.Role != model.RolePrimary && input.Role != model.RoleReference {
|
||||
return nil, "", ErrImageRole
|
||||
}
|
||||
if input.Role == model.RolePrimary {
|
||||
primaryCount++
|
||||
}
|
||||
if !utf8.ValidString(input.Note) || utf8.RuneCountInString(input.Note) > 500 {
|
||||
return nil, "", ErrInvalidNote
|
||||
}
|
||||
upload, exists := uploadByID[input.ClientID]
|
||||
if !exists {
|
||||
return nil, "", ErrImageMapping
|
||||
}
|
||||
orderedUploads = append(orderedUploads, upload)
|
||||
}
|
||||
if primaryCount != 1 {
|
||||
return nil, "", ErrImagePrimary
|
||||
}
|
||||
|
||||
validated, err := s.validateImages(orderedUploads)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
for index := range validated {
|
||||
validated[index].Position = inputs[index].Position
|
||||
validated[index].Role = inputs[index].Role
|
||||
validated[index].Note = strings.TrimSpace(inputs[index].Note)
|
||||
}
|
||||
return validated, roleRule, nil
|
||||
}
|
||||
|
||||
func (s *Service) submissionRoute(ctx context.Context, capability model.Capability) (corerouter.RouteSnapshot, model.PromptTemplate, error) {
|
||||
snapshot, err := s.routes.ActiveSnapshot(ctx, capability)
|
||||
if err != nil {
|
||||
return corerouter.RouteSnapshot{}, model.PromptTemplate{}, err
|
||||
}
|
||||
var promptTemplate model.PromptTemplate
|
||||
err = s.db.WithContext(ctx).
|
||||
Where("id=? AND template_key=? AND version=? AND capability=? AND enabled=TRUE", snapshot.PromptTemplateID, snapshot.PromptTemplateKey, snapshot.PromptTemplateVersion, capability).
|
||||
First(&promptTemplate).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return corerouter.RouteSnapshot{}, model.PromptTemplate{}, corerouter.ErrRouteUnavailable
|
||||
}
|
||||
return corerouter.RouteSnapshot{}, model.PromptTemplate{}, fmt.Errorf("load route prompt template: %w", err)
|
||||
}
|
||||
return snapshot, promptTemplate, nil
|
||||
}
|
||||
|
||||
func composeImagePrompt(base, defaultRoleRule, userRoleRule string, images []validatedImage) string {
|
||||
if len(images) == 0 {
|
||||
return base
|
||||
}
|
||||
rule := strings.TrimSpace(defaultRoleRule)
|
||||
if strings.TrimSpace(userRoleRule) != "" {
|
||||
rule = strings.TrimSpace(userRoleRule)
|
||||
}
|
||||
var result strings.Builder
|
||||
result.WriteString(base)
|
||||
if rule != "" {
|
||||
result.WriteString("\n\nImage role rule:\n")
|
||||
result.WriteString(rule)
|
||||
}
|
||||
result.WriteString("\n\nImage inputs (ordered by position):")
|
||||
for _, image := range images {
|
||||
fmt.Fprintf(&result, "\n- position=%d role=%s", image.Position, image.Role)
|
||||
if image.Note != "" {
|
||||
result.WriteString(" note=")
|
||||
result.WriteString(image.Note)
|
||||
}
|
||||
}
|
||||
return result.String()
|
||||
}
|
||||
@@ -7,7 +7,11 @@ import (
|
||||
"image"
|
||||
"image/color"
|
||||
"image/png"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.ilapage.cn/OPC/chorus/internal/core/model"
|
||||
)
|
||||
|
||||
func TestUploadValidationErrors(t *testing.T) {
|
||||
@@ -37,6 +41,74 @@ func TestUploadValidationErrors(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestImageMetadataValidationAndPromptComposition(t *testing.T) {
|
||||
service := &Service{config: Config{MaxPromptBytes: 1000, MaxImages: 3, MaxImageBytes: 1000, MaxUploadBytes: 3000, MaxImagePixels: 100}}
|
||||
metadata := ImageMetadata{
|
||||
Capability: model.CapabilityImageEdit,
|
||||
RoleRule: "user rule",
|
||||
Images: []ImageInput{
|
||||
{ClientID: "reference", Position: 1, Role: model.RoleReference, Note: "texture"},
|
||||
{ClientID: "primary", Position: 0, Role: model.RolePrimary, Note: "subject"},
|
||||
},
|
||||
}
|
||||
uploads := []Upload{
|
||||
{ClientID: "primary", Name: "a.png", DeclaredMIME: "image/png", Content: smallPNG(2, 2)},
|
||||
{ClientID: "reference", Name: "b.png", DeclaredMIME: "image/png", Content: smallPNG(2, 2)},
|
||||
}
|
||||
images, roleRule, err := service.validateImageSubmission(metadata, uploads)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if roleRule != "user rule" || len(images) != 2 || images[0].Position != 0 || images[0].Role != model.RolePrimary || images[1].Note != "texture" {
|
||||
t.Fatalf("validated images = %#v, roleRule=%q", images, roleRule)
|
||||
}
|
||||
rendered := composeImagePrompt("base", "default rule", roleRule, images)
|
||||
want := "base\n\nImage role rule:\nuser rule\n\nImage inputs (ordered by position):\n- position=0 role=primary note=subject\n- position=1 role=reference note=texture"
|
||||
if rendered != want {
|
||||
t.Fatalf("rendered = %q, want %q", rendered, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImageMetadataValidationErrors(t *testing.T) {
|
||||
service := &Service{config: Config{MaxPromptBytes: 8, MaxImages: 2, MaxImageBytes: 1000, MaxUploadBytes: 2000, MaxImagePixels: 100}}
|
||||
image := Upload{ClientID: "a", Name: "a.png", DeclaredMIME: "image/png", Content: smallPNG(2, 2)}
|
||||
valid := ImageMetadata{Capability: model.CapabilityImageEdit, Images: []ImageInput{{ClientID: "a", Position: 0, Role: model.RolePrimary}}}
|
||||
tests := []struct {
|
||||
name string
|
||||
metadata ImageMetadata
|
||||
uploads []Upload
|
||||
want error
|
||||
}{
|
||||
{"capability", ImageMetadata{Capability: model.CapabilityText}, nil, ErrInvalidCapability},
|
||||
{"generate files", ImageMetadata{Capability: model.CapabilityImageGenerate, Images: valid.Images}, []Upload{image}, ErrImageCount},
|
||||
{"missing files", valid, nil, ErrImageCount},
|
||||
{"client id", ImageMetadata{Capability: model.CapabilityImageEdit, Images: []ImageInput{{ClientID: "bad id", Role: model.RolePrimary}}}, []Upload{{ClientID: "bad id"}}, ErrInvalidClientID},
|
||||
{"position", ImageMetadata{Capability: model.CapabilityImageEdit, Images: []ImageInput{{ClientID: "a", Position: 1, Role: model.RolePrimary}}}, []Upload{image}, ErrImagePosition},
|
||||
{"role", ImageMetadata{Capability: model.CapabilityImageEdit, Images: []ImageInput{{ClientID: "a", Role: "other"}}}, []Upload{image}, ErrImageRole},
|
||||
{"primary", ImageMetadata{Capability: model.CapabilityImageEdit, Images: []ImageInput{{ClientID: "a", Role: model.RoleReference}}}, []Upload{image}, ErrImagePrimary},
|
||||
{"role rule", ImageMetadata{Capability: model.CapabilityImageEdit, RoleRule: "123456789", Images: valid.Images}, []Upload{image}, ErrInvalidRoleRule},
|
||||
{"note", ImageMetadata{Capability: model.CapabilityImageEdit, Images: []ImageInput{{ClientID: "a", Role: model.RolePrimary, Note: strings.Repeat("n", 501)}}}, []Upload{image}, ErrInvalidNote},
|
||||
{"mapping", valid, []Upload{{ClientID: "b"}}, ErrImageMapping},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
_, _, err := service.validateImageSubmission(test.metadata, test.uploads)
|
||||
if !errors.Is(err, test.want) {
|
||||
t.Fatalf("error=%v want=%v", err, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
invalidUTF8 := string([]byte{0xff})
|
||||
if utf8.ValidString(invalidUTF8) {
|
||||
t.Fatal("test string unexpectedly valid")
|
||||
}
|
||||
metadata := valid
|
||||
metadata.Images[0].Note = invalidUTF8
|
||||
if _, _, err := service.validateImageSubmission(metadata, []Upload{image}); !errors.Is(err, ErrInvalidNote) {
|
||||
t.Fatalf("invalid UTF-8 note error=%v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestImageCountAndTotalSize(t *testing.T) {
|
||||
service := &Service{config: Config{MaxImages: 1, MaxImageBytes: 1000, MaxUploadBytes: 100, MaxImagePixels: 100}}
|
||||
if _, err := service.validateImages(nil); !errors.Is(err, ErrImageCount) {
|
||||
|
||||
Reference in New Issue
Block a user