From 04953bd7df7d00f0a64504379fab36976779cf3c Mon Sep 17 00:00:00 2001 From: ila Date: Mon, 24 Aug 2026 15:35:38 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E4=BA=A4=E4=BB=98=20OpenAPI=20v1=20?= =?UTF-8?q?=E7=94=9F=E6=88=90=E6=8E=A5=E5=8F=A3=20(#43)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/02-architecture-and-code-map.md | 15 +- docs/03-business-rules-and-glossary.md | 14 +- docs/04-local-development-and-verification.md | 34 +- internal/core/apikey/repository.go | 16 +- portal/handler/mysql_integration_test.go | 197 ++++++++- portal/handler/openapi.go | 409 ++++++++++++++++++ portal/handler/router.go | 25 +- portal/openapi/openapi.json | 200 +++++++++ portal/openapi/spec.go | 10 + portal/openapi/spec_test.go | 50 +++ portal/service/idempotency.go | 74 ++++ portal/service/openapi.go | 65 +++ portal/service/service.go | 34 +- portal/service/validation_test.go | 36 ++ 14 files changed, 1159 insertions(+), 20 deletions(-) create mode 100644 portal/handler/openapi.go create mode 100644 portal/openapi/openapi.json create mode 100644 portal/openapi/spec.go create mode 100644 portal/openapi/spec_test.go create mode 100644 portal/service/idempotency.go create mode 100644 portal/service/openapi.go diff --git a/docs/02-architecture-and-code-map.md b/docs/02-architecture-and-code-map.md index 08884ce..3b4fbf7 100644 --- a/docs/02-architecture-and-code-map.md +++ b/docs/02-architecture-and-code-map.md @@ -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: 92ed37aff92d1f54fd18a8aac4f15835b8082251 -synchronized_at: 2026-08-24T06:51:38Z +wiki_revision: a660cd26350c8f7d242788e28c6ecc01bcf31f4d +synchronized_at: 2026-08-24T07:28:55Z # 架构与代码地图 @@ -405,3 +405,14 @@ Provider 限流检查发生在真实上游调用前。放行后才执行 `BeginP - `portal/web/templates/api_keys.html` 与 `static/api-keys.js` 实现加载、空、错误、停用、限流、一次展示、改名和不可恢复撤销状态;移动端改为卡片式行布局,桌面端保持紧凑表格。 - OpenAPI Bearer 认证和生成接口属于后续工单,不在 #42 中从 API Key 页面直接调用上游。 + + +## #43 OpenAPI v1 认证与生成接口 + +- `portal/handler` 将浏览器路由放在 session middleware 分组内,将 `/openapi/v1/*` 放在独立 API Key middleware 分组内;程序请求不会创建浏览器 Session Cookie,也不接受 Cookie、Query token 或管理员 JWT 回退。 +- API Key 认证解析固定 `chorus__` 格式,按 public id 定位后常量时间比较 secret hash,并在同一事务确认 Key 未撤销、未到期且 `users.status=active`。不存在、错误、撤销、到期和用户停用统一返回 `401 invalid_api_key`;成功使用时间最多每分钟写一次。 +- v1 提供文本与图片异步提交、HMAC cursor 历史、详情、输入、输出、缩略图和受认证的 `openapi.json`。全部响应带随机 `X-Request-ID` 并禁止缓存,资源 URL 固定指向 `/openapi/v1`。 +- 提交只读取必填 `Idempotency-Key` Header,限定 8 至 128 个可打印 ASCII 字符。相同用户跨浏览器/API 渠道重放相同请求返回 200;同键但 kind、Prompt、capability、图片角色元数据或文件内容不同返回 409。 +- 同步提交继续只做认证、严格校验、路由快照和事务落库,返回 202;没有 Provider 调用或 attempt。上游仍只由 worker 调用。 +- 仓库契约源为 `portal/openapi/openapi.json`,由 `portal/openapi` 嵌入二进制;合约测试固定 OpenAPI 3.1、八条路径、方法和 Bearer security。 + diff --git a/docs/03-business-rules-and-glossary.md b/docs/03-business-rules-and-glossary.md index c1dba57..401c758 100644 --- a/docs/03-business-rules-and-glossary.md +++ b/docs/03-business-rules-and-glossary.md @@ -2,8 +2,8 @@ generated: true (请先修改 Gitea Wiki,禁止直接编辑本文件) wiki_page: Business-Rules-and-Glossary wiki_url: https://git.ilapage.cn/OPC/chorus/wiki/Business-Rules-and-Glossary.- -wiki_revision: 1bdd0adffefbf1fbb33ddfdab777e2f27e6aef83 -synchronized_at: 2026-08-24T06:51:46Z +wiki_revision: e63946a4f4b0ca4acbad0a0ba0e75b9ed70db95a +synchronized_at: 2026-08-24T07:29:01Z # 业务规则与术语 @@ -156,3 +156,13 @@ synchronized_at: 2026-08-24T06:51:46Z - 撤销立即生效、不能恢复;重复撤销保持幂等。已过期或已撤销 Key 仍保留名称、前缀和状态,便于用户识别历史记录。 - API Key 生命周期不包含注册、管理员代管、计费、点数、每日配额或 OpenAPI 调用;请求限流与审计分别由后续单元工单交付。 + + +## OpenAPI v1 调用规则(#43) + +- 程序接口只接受单个 `Authorization: Bearer ` Header。Cookie、Query token、管理员 JWT 和浏览器 CSRF 不能认证 OpenAPI;API Key 也不能认证浏览器 `/api/*` 或管理员接口。 +- 请求成功或失败均返回 `X-Request-ID`,并使用 `{error:{code,message}}` 错误信封。OpenAPI 响应禁止缓存;不存在、错误、撤销、到期和停用用户的 Key 不能通过响应内容区分。 +- 文本和图片提交必须提供单个 `Idempotency-Key` Header,去除首尾空格前后必须一致,长度 8 至 128 字节且只能包含可打印 ASCII。正文或 multipart 字段中的同名值不被接受。 +- 同一用户的幂等键跨浏览器和 OpenAPI 共用。只有完整请求语义一致才是 200 重放;图片会比较 capability、Prompt、role rule、输入顺序/角色/备注、文件元数据和文件内容。任何差异返回 409,且不会覆盖原任务或写第二份输入。 +- 新提交立即返回 202 pending,程序需要用历史或详情接口查询状态;OpenAPI 不提供同步等待、流式、取消、批量、webhook、Swagger UI 或 SDK。 + diff --git a/docs/04-local-development-and-verification.md b/docs/04-local-development-and-verification.md index 081b46f..efa0e50 100644 --- a/docs/04-local-development-and-verification.md +++ b/docs/04-local-development-and-verification.md @@ -2,8 +2,8 @@ generated: true (请先修改 Gitea Wiki,禁止直接编辑本文件) wiki_page: Local-Development-and-Verification wiki_url: https://git.ilapage.cn/OPC/chorus/wiki/Local-Development-and-Verification.- -wiki_revision: 71cf8fd56b2a18494d380464f59fc58d81283bd9 -synchronized_at: 2026-08-24T06:51:53Z +wiki_revision: e8d57350c40ad6a6ac996b5a3e5903ca782a1d1c +synchronized_at: 2026-08-24T07:29:07Z # 本地开发与验证 @@ -441,3 +441,33 @@ pnpm --dir portal/web build 浏览器 E2E 使用 `portal/web/e2e/fixture` 的合成用户和独立库,在 375、768、1024、1440 四种视口执行创建、一次展示、关闭后清除、刷新不可恢复、改名和撤销,并检查无外部请求、无页面横向溢出和至少 44px 的操作目标。测试截图只能在完整 token 已从 DOM 清除后生成。 + + +## 使用和验证 OpenAPI v1(#43) + +完整契约保存在 `portal/openapi/openapi.json`。运行中的同版本接口在 `/openapi/v1/openapi.json` 返回该文件;该路径同样要求有效 API Key。不要把真实 Key 写入脚本、命令历史、工单或文档,可在受保护的当前终端临时设置并在完成后清除: + +```powershell +$env:CHORUS_API_KEY = "<本次创建且已妥善保存的 API Key>" +$headers = @{ + Authorization = "Bearer $env:CHORUS_API_KEY" + "Idempotency-Key" = [guid]::NewGuid().ToString() +} +$body = @{ prompt = "生成一段简短说明" } | ConvertTo-Json +Invoke-RestMethod -Method Post -Uri "http://127.0.0.1:8080/openapi/v1/generations/text" ` + -Headers $headers -ContentType "application/json" -Body $body +$env:CHORUS_API_KEY = $null +``` + +首次相同请求返回 202;用完全相同的 Header 和正文重放返回 200 和同一 generation id。相同幂等键配不同请求返回 409。图片接口使用 multipart 的 `prompt`、`metadata` 和 `files[]`;`metadata` 必须包含 capability、role_rule、images,文件声明与 part 一一对应。详情中的文件 URL 已指向 `/openapi/v1`,下载仍需同一个 Bearer Header。 + +普通及契约回归: + +```powershell +go test ./portal/openapi ./portal/service ./portal/handler +go vet ./portal/openapi ./portal/service ./portal/handler +go test ./... +``` + +真实 Handler 集成测试必须使用显式命名、可丢弃的 MySQL 8 隔离库。#43 已在 `chorus_mvp2_openapi_test` 覆盖三条认证链隔离、无效/到期/撤销/停用 Key、JSON/multipart、跨渠道幂等与冲突、游标、跨用户详情和文件访问;测试结束已删除该库,未调用真实 Provider。 + diff --git a/internal/core/apikey/repository.go b/internal/core/apikey/repository.go index c2e6607..70743e1 100644 --- a/internal/core/apikey/repository.go +++ b/internal/core/apikey/repository.go @@ -10,6 +10,7 @@ import ( "git.ilapage.cn/OPC/chorus/internal/core/model" "gorm.io/gorm" + "gorm.io/gorm/clause" ) var ( @@ -20,6 +21,7 @@ var ( type Repository interface { Create(ctx context.Context, key *model.APIKey) error ByPublicID(ctx context.Context, publicID string) (model.APIKey, error) + ByPublicIDLocked(ctx context.Context, publicID string) (model.APIKey, error) ByIDForUser(ctx context.Context, id, userID uint64) (model.APIKey, error) ListForUser(ctx context.Context, userID uint64) ([]model.APIKey, error) Rename(ctx context.Context, id, userID uint64, name string) (bool, error) @@ -52,12 +54,24 @@ func (r *GORMRepository) Create(ctx context.Context, key *model.APIKey) error { } func (r *GORMRepository) ByPublicID(ctx context.Context, publicID string) (model.APIKey, error) { + return r.byPublicID(ctx, publicID, false) +} + +func (r *GORMRepository) ByPublicIDLocked(ctx context.Context, publicID string) (model.APIKey, error) { + return r.byPublicID(ctx, publicID, true) +} + +func (r *GORMRepository) byPublicID(ctx context.Context, publicID string, lock bool) (model.APIKey, error) { publicID = strings.TrimSpace(publicID) if len(publicID) != 24 { return model.APIKey{}, ErrInvalidAPIKey } var key model.APIKey - if err := r.db.WithContext(ctx).Where("public_id = ?", publicID).First(&key).Error; err != nil { + query := r.db.WithContext(ctx) + if lock { + query = query.Clauses(clause.Locking{Strength: "SHARE"}) + } + if err := query.Where("public_id = ?", publicID).First(&key).Error; err != nil { return model.APIKey{}, mapNotFound(err, "find api key by public id") } return key, nil diff --git a/portal/handler/mysql_integration_test.go b/portal/handler/mysql_integration_test.go index 979cb07..3400825 100644 --- a/portal/handler/mysql_integration_test.go +++ b/portal/handler/mysql_integration_test.go @@ -244,6 +244,80 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { if err != nil || parsedKey.PublicID != storedKey.PublicID || !bytes.Equal(parsedKey.SecretHash, storedKey.SecretHash) { t.Fatalf("stored API key does not match one-time token: %v", err) } + openAPIKeyA := storeTestAPIKey(t, db, users[0].ID, "OpenAPI A", nil, nil) + openAPIKeyB := storeTestAPIKey(t, db, users[1].ID, "OpenAPI B", nil, nil) + expiredAt := time.Now().UTC().Add(-time.Minute) + expiredKey := storeTestAPIKey(t, db, users[0].ID, "Expired OpenAPI", &expiredAt, nil) + missingAuth := openAPIRequest(router, "", http.MethodGet, "/openapi/v1/openapi.json", nil, "", "") + expiredAuth := openAPIRequest(router, expiredKey.Token, http.MethodGet, "/openapi/v1/openapi.json", nil, "", "") + invalidAuth := openAPIRequest(router, "administrator-jwt", http.MethodGet, "/openapi/v1/openapi.json", nil, "", "") + queryAuth := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/openapi.json?Authorization=ignored", nil, "", "") + for name, response := range map[string]*httptest.ResponseRecorder{"missing": missingAuth, "expired": expiredAuth, "invalid": invalidAuth, "query": queryAuth} { + if response.Code != http.StatusUnauthorized || response.Body.String() != missingAuth.Body.String() || response.Header().Get("WWW-Authenticate") != "Bearer" || response.Header().Get("Cache-Control") != "no-store" || response.Header().Get("X-Request-ID") == "" || response.Header().Get("Set-Cookie") != "" { + t.Fatalf("%s OpenAPI auth=%d headers=%v body=%s", name, response.Code, response.Header(), response.Body.String()) + } + } + browserCookieRequest := httptest.NewRequest(http.MethodGet, "/openapi/v1/openapi.json", nil) + browserCookieRequest.Header.Set("Authorization", "Bearer "+openAPIKeyA.Token) + browserCookieRequest.AddCookie(clientA.cookie) + browserCookieResponse := httptest.NewRecorder() + router.ServeHTTP(browserCookieResponse, browserCookieRequest) + if browserCookieResponse.Code != http.StatusUnauthorized { + t.Fatalf("OpenAPI accepted browser cookie=%d", browserCookieResponse.Code) + } + specResponse := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/openapi.json", nil, "", "") + var specification struct { + OpenAPI string `json:"openapi"` + } + if specResponse.Code != http.StatusOK || json.Unmarshal(specResponse.Body.Bytes(), &specification) != nil || specification.OpenAPI != "3.1.0" || specResponse.Header().Get("Set-Cookie") != "" || specResponse.Header().Get("X-Request-ID") == "" { + t.Fatalf("OpenAPI specification=%d headers=%v body=%s", specResponse.Code, specResponse.Header(), specResponse.Body.String()) + } + var usedKey model.APIKey + if err := db.Where("public_id = ?", openAPIKeyA.PublicID).First(&usedKey).Error; err != nil || usedKey.LastUsedAt == nil { + t.Fatalf("OpenAPI key last_used_at was not updated: %#v %v", usedKey, err) + } + browserOnly := clientA.do(http.MethodGet, "/openapi/v1/generations", nil, "") + if browserOnly.Code != http.StatusUnauthorized { + t.Fatalf("OpenAPI accepted browser session=%d", browserOnly.Code) + } + apiKeyOnBrowserRoute := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/api/generations", nil, "", "") + if apiKeyOnBrowserRoute.Code != http.StatusUnauthorized { + t.Fatalf("browser route accepted API key=%d", apiKeyOnBrowserRoute.Code) + } + + openAPITextKey := "openapi-text-" + suffix + openAPITextBody, _ := json.Marshal(map[string]string{"prompt": "OpenAPI text prompt"}) + openAPIText := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", openAPITextKey) + if openAPIText.Code != http.StatusAccepted { + t.Fatalf("OpenAPI text submit=%d %s", openAPIText.Code, openAPIText.Body.String()) + } + openAPITextID := responseID(t, openAPIText) + openAPITextReplay := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", openAPITextKey) + if openAPITextReplay.Code != http.StatusOK || responseID(t, openAPITextReplay) != openAPITextID { + t.Fatalf("OpenAPI text replay=%d %s", openAPITextReplay.Code, openAPITextReplay.Body.String()) + } + conflictBody, _ := json.Marshal(map[string]string{"prompt": "different OpenAPI prompt"}) + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", conflictBody, "application/json", openAPITextKey); response.Code != http.StatusConflict || !strings.Contains(response.Body.String(), `"code":"idempotency_conflict"`) { + t.Fatalf("OpenAPI text conflict=%d %s", response.Code, response.Body.String()) + } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", ""); response.Code != http.StatusBadRequest { + t.Fatalf("missing Idempotency-Key=%d %s", response.Code, response.Body.String()) + } + unknownFieldBody, _ := json.Marshal(map[string]string{"prompt": "prompt", "idempotency_key": "body-key"}) + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", unknownFieldBody, "application/json", "header-key-123"); response.Code != http.StatusBadRequest { + t.Fatalf("body idempotency key was accepted=%d %s", response.Code, response.Body.String()) + } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/text", openAPITextBody, "application/json", "非ASCII幂等键"); response.Code != http.StatusBadRequest { + t.Fatalf("non-ASCII Idempotency-Key=%d %s", response.Code, response.Body.String()) + } + openAPICrossChannelBody, _ := json.Marshal(map[string]string{"idempotency_key": openAPITextKey, "prompt": "OpenAPI text prompt"}) + if response := clientA.do(http.MethodPost, "/api/generations/text", openAPICrossChannelBody, "application/json"); response.Code != http.StatusOK || responseID(t, response) != openAPITextID { + t.Fatalf("cross-channel idempotent replay=%d %s", response.Code, response.Body.String()) + } + var openAPIGeneration model.Generation + if err := db.First(&openAPIGeneration, openAPITextID).Error; err != nil || openAPIGeneration.Status != model.StatusPending || openAPIGeneration.AttemptCount != 0 || openAPIGeneration.ProviderModelID != nil { + t.Fatalf("OpenAPI synchronous submit touched worker fields: %#v %v", openAPIGeneration, err) + } listKeys := clientA.do(http.MethodGet, "/api/api-keys", nil, "") if listKeys.Code != http.StatusOK || !strings.Contains(listKeys.Body.String(), "自动化脚本") || strings.Contains(listKeys.Body.String(), createdKey.Token) || strings.Contains(listKeys.Body.String(), `secret_hash`) { t.Fatalf("list API keys=%d %s", listKeys.Code, listKeys.Body.String()) @@ -307,7 +381,8 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { } imageKey := "image-key-" + suffix - 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) + imageParts := []uploadPart{{clientID: "primary-a", name: "first.png", mime: "image/png", data: testPNG(4, 4), position: 0, role: model.RolePrimary, note: "keep edges"}} + imageResponse := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "user role rule", imageParts), lastMultipartType) if imageResponse.Code != 202 { t.Fatalf("image submit=%d %s", imageResponse.Code, imageResponse.Body.String()) } @@ -323,10 +398,14 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { 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) + imageReplay := clientA.do(http.MethodPost, "/api/generations/image", multipartBody(t, imageKey, "edit image", model.CapabilityImageEdit, "user role rule", imageParts), lastMultipartType) if imageReplay.Code != 200 || responseID(t, imageReplay) != imageID { t.Fatalf("image replay=%d %s", imageReplay.Code, imageReplay.Body.String()) } + imageConflict := 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 imageConflict.Code != http.StatusConflict || !strings.Contains(imageConflict.Body.String(), `"code":"idempotency_conflict"`) { + t.Fatalf("image conflict=%d %s", imageConflict.Code, imageConflict.Body.String()) + } if err := db.Table("route_pools").Where("id=?", routeIDs[2]).Update("enabled", true).Error; err != nil { t.Fatal(err) } @@ -335,6 +414,25 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { if inputCount != 1 { t.Fatalf("replay wrote %d inputs", inputCount) } + openAPIImageKey := "openapi-image-" + suffix + openAPIImageParts := []uploadPart{{clientID: "openapi-primary", name: "openapi.png", mime: "image/png", data: testPNG(3, 3), position: 0, role: model.RolePrimary, note: "preserve subject"}} + openAPIImageBody, openAPIImageType := openAPIMultipartBody(t, "OpenAPI image edit", model.CapabilityImageEdit, "OpenAPI role rule", openAPIImageParts) + openAPIImage := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/image", openAPIImageBody, openAPIImageType, openAPIImageKey) + if openAPIImage.Code != http.StatusAccepted { + t.Fatalf("OpenAPI image submit=%d %s", openAPIImage.Code, openAPIImage.Body.String()) + } + openAPIImageID := responseID(t, openAPIImage) + openAPIImageBody, openAPIImageType = openAPIMultipartBody(t, "OpenAPI image edit", model.CapabilityImageEdit, "OpenAPI role rule", openAPIImageParts) + openAPIImageReplay := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/image", openAPIImageBody, openAPIImageType, openAPIImageKey) + if openAPIImageReplay.Code != http.StatusOK || responseID(t, openAPIImageReplay) != openAPIImageID { + t.Fatalf("OpenAPI image replay=%d %s", openAPIImageReplay.Code, openAPIImageReplay.Body.String()) + } + changedOpenAPIImageParts := append([]uploadPart(nil), openAPIImageParts...) + changedOpenAPIImageParts[0].data = testPNG(5, 5) + openAPIImageBody, openAPIImageType = openAPIMultipartBody(t, "OpenAPI image edit", model.CapabilityImageEdit, "OpenAPI role rule", changedOpenAPIImageParts) + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodPost, "/openapi/v1/generations/image", openAPIImageBody, openAPIImageType, openAPIImageKey); response.Code != http.StatusConflict || !strings.Contains(response.Body.String(), `"code":"idempotency_conflict"`) { + t.Fatalf("OpenAPI image conflict=%d %s", response.Code, response.Body.String()) + } 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()) @@ -376,6 +474,23 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { 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()) } + openAPIHistory := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/generations?limit=2", nil, "", "") + var openAPIHistoryPage struct { + Items []struct { + ID uint64 `json:"id"` + } `json:"items"` + NextCursor string `json:"next_cursor"` + HasMore bool `json:"has_more"` + } + if openAPIHistory.Code != http.StatusOK || json.Unmarshal(openAPIHistory.Body.Bytes(), &openAPIHistoryPage) != nil || len(openAPIHistoryPage.Items) != 2 || !openAPIHistoryPage.HasMore || openAPIHistoryPage.NextCursor == "" { + t.Fatalf("OpenAPI history=%d %s", openAPIHistory.Code, openAPIHistory.Body.String()) + } + if response := openAPIRequest(router, openAPIKeyB.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d", openAPITextID), nil, "", ""); response.Code != http.StatusNotFound { + t.Fatalf("cross-user OpenAPI detail=%d %s", response.Code, response.Body.String()) + } + if response := openAPIRequest(router, openAPIKeyB.Token, http.MethodGet, "/openapi/v1/generations?limit=2&cursor="+openAPIHistoryPage.NextCursor, nil, "", ""); response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), `"code":"invalid_cursor"`) { + t.Fatalf("cross-user OpenAPI cursor=%d %s", response.Code, response.Body.String()) + } clientB := &apiClient{router: router} clientB.start(t) @@ -402,6 +517,9 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { if response := clientA.do(http.MethodGet, "/api/api-keys", nil, ""); response.Code != http.StatusForbidden || !strings.Contains(response.Body.String(), `"code":"account_disabled"`) { t.Fatalf("disabled user API keys=%d %s", response.Code, response.Body.String()) } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/generations", nil, "", ""); response.Code != http.StatusUnauthorized || !strings.Contains(response.Body.String(), `"code":"invalid_api_key"`) { + t.Fatalf("disabled user OpenAPI=%d %s", response.Code, response.Body.String()) + } if err := db.Model(&model.User{}).Where("id = ?", users[0].ID).Update("status", "active").Error; err != nil { t.Fatal(err) } @@ -449,6 +567,29 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { if response := clientA.do(http.MethodGet, fmt.Sprintf("/api/generations/%d/outputs/%d/thumbnail", imageID, output.ID), nil, ""); response.Code != 200 { t.Fatalf("own thumbnail=%d", response.Code) } + openAPIDetail := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d", imageID), nil, "", "") + if openAPIDetail.Code != http.StatusOK || !strings.Contains(openAPIDetail.Body.String(), fmt.Sprintf("/openapi/v1/generations/%d/inputs/%d", imageID, inputs[0].ID)) || !strings.Contains(openAPIDetail.Body.String(), fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", imageID, output.ID)) || strings.Contains(openAPIDetail.Body.String(), storageRoot) || strings.Contains(openAPIDetail.Body.String(), original.Key) { + t.Fatalf("OpenAPI detail=%d %s", openAPIDetail.Code, openAPIDetail.Body.String()) + } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/inputs/%d", imageID, inputs[0].ID), nil, "", ""); response.Code != http.StatusOK || response.Header().Get("Content-Type") != "image/png" { + t.Fatalf("OpenAPI input=%d %s", response.Code, response.Header().Get("Content-Type")) + } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", imageID, output.ID), nil, "", ""); response.Code != http.StatusOK || response.Header().Get("Content-Type") != "image/png" { + t.Fatalf("OpenAPI output=%d %s", response.Code, response.Header().Get("Content-Type")) + } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d/thumbnail", imageID, output.ID), nil, "", ""); response.Code != http.StatusOK || response.Header().Get("Content-Type") != "image/png" { + t.Fatalf("OpenAPI thumbnail=%d %s", response.Code, response.Header().Get("Content-Type")) + } + if response := openAPIRequest(router, openAPIKeyB.Token, http.MethodGet, fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", imageID, output.ID), nil, "", ""); response.Code != http.StatusNotFound { + t.Fatalf("cross-user OpenAPI output=%d %s", response.Code, response.Body.String()) + } + revokedAt := time.Now().UTC() + if err := db.Model(&model.APIKey{}).Where("public_id = ?", openAPIKeyA.PublicID).Update("revoked_at", revokedAt).Error; err != nil { + t.Fatal(err) + } + if response := openAPIRequest(router, openAPIKeyA.Token, http.MethodGet, "/openapi/v1/generations", nil, "", ""); response.Code != http.StatusUnauthorized || response.Body.String() != missingAuth.Body.String() { + t.Fatalf("revoked OpenAPI key=%d %s", response.Code, response.Body.String()) + } anonymous := &apiClient{router: router} anonymous.start(t) @@ -486,6 +627,58 @@ type uploadPart struct { var lastMultipartType string +func openAPIRequest(router http.Handler, token, method, path string, body []byte, contentType, idempotencyKey string) *httptest.ResponseRecorder { + request := httptest.NewRequest(method, path, bytes.NewReader(body)) + request.RemoteAddr = "192.0.2.20:1234" + if token != "" { + request.Header.Set("Authorization", "Bearer "+token) + } + if contentType != "" { + request.Header.Set("Content-Type", contentType) + } + if idempotencyKey != "" { + request.Header.Set("Idempotency-Key", idempotencyKey) + } + response := httptest.NewRecorder() + router.ServeHTTP(response, request) + return response +} + +func storeTestAPIKey(t *testing.T, db *gorm.DB, userID uint64, name string, expiresAt, revokedAt *time.Time) platformapikey.Credential { + t.Helper() + credential, err := platformapikey.Generate() + if err != nil { + t.Fatal(err) + } + row := model.APIKey{UserID: userID, Name: name, PublicID: credential.PublicID, KeyPrefix: credential.KeyPrefix, SecretHash: credential.SecretHash, ExpiresAt: expiresAt, RevokedAt: revokedAt} + if err := db.Create(&row).Error; err != nil { + t.Fatal(err) + } + return credential +} + +func openAPIMultipartBody(t *testing.T, prompt string, capability model.Capability, roleRule string, parts []uploadPart) ([]byte, string) { + t.Helper() + var body bytes.Buffer + writer := multipart.NewWriter(&body) + _ = 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="files[%s]"; filename="%s"`, item.clientID, item.name)) + header.Set("Content-Type", item.mime) + part, _ := writer.CreatePart(header) + _, _ = part.Write(item.data) + } + _ = writer.Close() + return body.Bytes(), writer.FormDataContentType() +} + func multipartBody(t *testing.T, key, prompt string, capability model.Capability, roleRule string, parts []uploadPart) []byte { t.Helper() var body bytes.Buffer diff --git a/portal/handler/openapi.go b/portal/handler/openapi.go new file mode 100644 index 0000000..3923e68 --- /dev/null +++ b/portal/handler/openapi.go @@ -0,0 +1,409 @@ +package handler + +import ( + "crypto/rand" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "mime" + "net/http" + "strconv" + "strings" + "sync/atomic" + "time" + + "git.ilapage.cn/OPC/chorus/internal/core/model" + "git.ilapage.cn/OPC/chorus/portal/openapi" + "git.ilapage.cn/OPC/chorus/portal/service" + "github.com/gin-gonic/gin" +) + +const openAPIPrincipalContextKey = "openapi_principal" + +var openAPIRequestIDCounter atomic.Uint64 + +func (h *Handler) openAPIHeaders(c *gin.Context) { + c.Header("X-Request-ID", newOpenAPIRequestID()) + noStore(c) + c.Next() +} + +func newOpenAPIRequestID() string { + requestID := make([]byte, 16) + if _, err := rand.Read(requestID); err == nil { + return hex.EncodeToString(requestID) + } + return fmt.Sprintf("%x-%x", time.Now().UTC().UnixNano(), openAPIRequestIDCounter.Add(1)) +} + +func (h *Handler) requireAPIKey(c *gin.Context) { + if c.GetHeader("Cookie") != "" || hasCredentialQuery(c) { + h.invalidAPIKey(c) + return + } + headers := c.Request.Header.Values("Authorization") + if len(headers) != 1 { + h.invalidAPIKey(c) + return + } + parts := strings.Fields(headers[0]) + if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") { + h.invalidAPIKey(c) + return + } + principal, err := h.service.AuthenticateAPIKey(c.Request.Context(), parts[1], time.Now().UTC()) + if err != nil { + if errors.Is(err, service.ErrInvalidAPIKeyCredential) { + h.invalidAPIKey(c) + return + } + writeError(c, http.StatusInternalServerError, "internal_error", "request could not be completed") + c.Abort() + return + } + c.Set(openAPIPrincipalContextKey, principal) + c.Next() +} + +func (h *Handler) invalidAPIKey(c *gin.Context) { + c.Header("WWW-Authenticate", "Bearer") + writeError(c, http.StatusUnauthorized, "invalid_api_key", "API key is missing or invalid") + c.Abort() +} + +func hasCredentialQuery(c *gin.Context) bool { + for name := range c.Request.URL.Query() { + for _, credentialName := range []string{"api_key", "access_token", "token", "authorization"} { + if strings.EqualFold(name, credentialName) { + return true + } + } + } + return false +} + +func currentAPIPrincipal(c *gin.Context) service.APIPrincipal { + value, _ := c.Get(openAPIPrincipalContextKey) + principal, _ := value.(service.APIPrincipal) + return principal +} + +func (h *Handler) openAPISpec(c *gin.Context) { + if !requireNoQuery(c) { + return + } + c.Data(http.StatusOK, "application/json; charset=utf-8", openapi.Spec()) +} + +func (h *Handler) openAPISubmitText(c *gin.Context) { + if !requireNoQuery(c) { + return + } + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 64<<10) + var input struct { + Prompt string `json:"prompt"` + } + if !decodeOpenAPIJSON(c, &input) { + return + } + idempotencyKey, ok := openAPIIdempotencyKey(c) + if !ok { + return + } + principal := currentAPIPrincipal(c) + generation, created, err := h.service.SubmitText(c.Request.Context(), principal.UserID, idempotencyKey, input.Prompt) + if err != nil { + h.serviceError(c, err) + return + } + status := http.StatusAccepted + if !created { + status = http.StatusOK + } + c.JSON(status, openAPIGenerationResponse(generation)) +} + +func decodeOpenAPIJSON(c *gin.Context, target any) bool { + if mediaType, _, err := mime.ParseMediaType(c.GetHeader("Content-Type")); err != nil || mediaType != "application/json" { + writeError(c, http.StatusBadRequest, "invalid_request", "request body is invalid") + return false + } + decoder := json.NewDecoder(c.Request.Body) + decoder.DisallowUnknownFields() + if err := decoder.Decode(target); err != nil || decoder.Decode(&struct{}{}) != io.EOF { + writeError(c, http.StatusBadRequest, "invalid_request", "request body is invalid") + return false + } + return true +} + +func (h *Handler) openAPISubmitImages(c *gin.Context) { + if !requireNoQuery(c) { + return + } + mediaType, _, err := mime.ParseMediaType(c.GetHeader("Content-Type")) + if err != nil || mediaType != "multipart/form-data" { + writeError(c, http.StatusBadRequest, "invalid_request", "request body is invalid") + return + } + c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, h.maxUploadBytes+(1<<20)) + if err := c.Request.ParseMultipartForm(h.maxUploadBytes); err != nil { + writeError(c, http.StatusBadRequest, "upload_too_large", "upload exceeds configured limit") + return + } + defer c.Request.MultipartForm.RemoveAll() + if !exactMultipartValues(c, "prompt", "metadata") { + writeError(c, http.StatusBadRequest, "invalid_request", "multipart fields are invalid") + return + } + var input struct { + Capability *model.Capability `json:"capability"` + RoleRule *string `json:"role_rule"` + Images *[]service.ImageInput `json:"images"` + } + decoder := json.NewDecoder(strings.NewReader(c.Request.MultipartForm.Value["metadata"][0])) + decoder.DisallowUnknownFields() + if err := decoder.Decode(&input); err != nil || decoder.Decode(&struct{}{}) != io.EOF || input.Capability == nil || input.RoleRule == nil || input.Images == nil { + writeError(c, http.StatusBadRequest, "invalid_metadata", "image metadata is invalid") + return + } + metadata := service.ImageMetadata{Capability: *input.Capability, RoleRule: *input.RoleRule, Images: *input.Images} + uploads, ok := h.openAPIUploads(c, metadata) + if !ok { + return + } + idempotencyKey, ok := openAPIIdempotencyKey(c) + if !ok { + return + } + principal := currentAPIPrincipal(c) + generation, created, submitErr := h.service.SubmitImages(c.Request.Context(), principal.UserID, idempotencyKey, c.Request.MultipartForm.Value["prompt"][0], metadata, uploads) + if submitErr != nil { + h.serviceError(c, submitErr) + return + } + status := http.StatusAccepted + if !created { + status = http.StatusOK + } + c.JSON(status, openAPIGenerationResponse(generation)) +} + +func exactMultipartValues(c *gin.Context, required ...string) bool { + if len(c.Request.MultipartForm.Value) != len(required) { + return false + } + for _, name := range required { + if values := c.Request.MultipartForm.Value[name]; len(values) != 1 { + return false + } + } + return true +} + +func (h *Handler) openAPIUploads(c *gin.Context, metadata service.ImageMetadata) ([]service.Upload, bool) { + expectedFields := make(map[string]string, len(metadata.Images)) + for _, item := range metadata.Images { + field := "files[" + item.ClientID + "]" + if _, exists := expectedFields[field]; exists { + writeError(c, http.StatusBadRequest, "invalid_image_mapping", "image files do not match metadata") + return nil, false + } + expectedFields[field] = item.ClientID + } + if len(c.Request.MultipartForm.File) != len(expectedFields) { + writeError(c, http.StatusBadRequest, "invalid_image_mapping", "image files do not match metadata") + return nil, false + } + uploads := make([]service.Upload, 0, len(expectedFields)) + for field, clientID := range expectedFields { + headers := c.Request.MultipartForm.File[field] + if len(headers) != 1 { + writeError(c, http.StatusBadRequest, "invalid_image_mapping", "image files do not match metadata") + return nil, false + } + header := headers[0] + file, err := header.Open() + if err != nil { + writeFileError(c, int(metadataPosition(metadata, clientID)), safeName(header.Filename), "invalid_image_content") + return nil, false + } + content, readErr := io.ReadAll(io.LimitReader(file, h.maxUploadBytes+1)) + closeErr := file.Close() + if readErr != nil || closeErr != nil { + writeFileError(c, int(metadataPosition(metadata, clientID)), safeName(header.Filename), "invalid_image_content") + return nil, false + } + uploads = append(uploads, service.Upload{ClientID: clientID, Name: header.Filename, DeclaredMIME: header.Header.Get("Content-Type"), Content: content}) + } + return uploads, true +} + +func (h *Handler) openAPIHistory(c *gin.Context) { + if !queryNamesAllowed(c, "limit", "cursor") { + writeError(c, http.StatusBadRequest, "invalid_request", "query parameters are invalid") + return + } + limit := h.service.DefaultHistoryLimit() + if values, exists := c.GetQueryArray("limit"); exists { + if len(values) != 1 { + writeError(c, http.StatusBadRequest, "invalid_limit", "history limit is invalid") + return + } + parsed, parseErr := strconv.Atoi(values[0]) + if parseErr != nil { + writeError(c, http.StatusBadRequest, "invalid_limit", "history limit is invalid") + return + } + limit = parsed + } + cursor := "" + if values, exists := c.GetQueryArray("cursor"); exists { + if len(values) != 1 || values[0] == "" { + writeError(c, http.StatusBadRequest, "invalid_cursor", "history cursor is invalid") + return + } + cursor = values[0] + } + page, err := h.service.History(c.Request.Context(), currentAPIPrincipal(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, openAPIGenerationResponse(row)) + } + c.JSON(http.StatusOK, gin.H{"items": items, "next_cursor": page.NextCursor, "has_more": page.HasMore}) +} + +func queryNamesAllowed(c *gin.Context, allowed ...string) bool { + set := make(map[string]struct{}, len(allowed)) + for _, name := range allowed { + set[name] = struct{}{} + } + for name := range c.Request.URL.Query() { + if _, ok := set[name]; !ok { + return false + } + } + return true +} + +func requireNoQuery(c *gin.Context) bool { + if queryNamesAllowed(c) { + return true + } + writeError(c, http.StatusBadRequest, "invalid_request", "query parameters are invalid") + return false +} + +func openAPIIdempotencyKey(c *gin.Context) (string, bool) { + values := c.Request.Header.Values("Idempotency-Key") + if len(values) != 1 { + writeError(c, http.StatusBadRequest, "invalid_idempotency_key", "idempotency key is invalid") + return "", false + } + return values[0], true +} + +func (h *Handler) openAPIDetail(c *gin.Context) { + if !queryNamesAllowed(c) { + writeError(c, http.StatusBadRequest, "invalid_request", "query parameters are invalid") + return + } + id, ok := uintParam(c, "id") + if !ok { + return + } + detail, err := h.service.ByID(c.Request.Context(), currentAPIPrincipal(c).UserID, id) + if err != nil { + h.serviceError(c, err) + return + } + response := openAPIGenerationResponse(detail.Generation) + 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, "note": input.Note, "name": input.OriginalName, "mime_type": input.MIMEType, "size_bytes": input.SizeBytes, "width": input.Width, "height": input.Height, "url": fmt.Sprintf("/openapi/v1/generations/%d/inputs/%d", id, input.ID)}) + } + outputs := make([]gin.H, 0, len(detail.Outputs)) + for _, output := range detail.Outputs { + item := gin.H{"id": output.ID, "kind": output.Kind, "mime_type": output.MIMEType, "size_bytes": output.SizeBytes, "width": output.Width, "height": output.Height, "created_at": output.CreatedAt.UTC()} + if output.TextContent != nil { + item["text"] = *output.TextContent + } + if output.StorageKey != nil { + item["url"] = fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d", id, output.ID) + if output.ThumbnailStorageKey != nil { + item["thumbnail_url"] = fmt.Sprintf("/openapi/v1/generations/%d/outputs/%d/thumbnail", id, output.ID) + } + } + outputs = append(outputs, item) + } + response["role_rule"] = detail.Generation.RoleRule + response["inputs"] = inputs + response["outputs"] = outputs + c.JSON(http.StatusOK, response) +} + +func (h *Handler) openAPIInput(c *gin.Context) { + generationID, inputID, ok := openAPIFileIDs(c, "inputID") + if !ok { + return + } + reader, object, name, err := h.service.OpenInput(c.Request.Context(), currentAPIPrincipal(c).UserID, generationID, inputID) + if err != nil { + h.serviceError(c, err) + return + } + defer reader.Close() + c.Header("Content-Type", object.ContentType) + c.Header("Content-Disposition", mime.FormatMediaType("inline", map[string]string{"filename": name})) + c.Status(http.StatusOK) + _, _ = io.Copy(c.Writer, reader) +} + +func (h *Handler) openAPIOutput(c *gin.Context) { h.writeOpenAPIOutput(c, false) } +func (h *Handler) openAPIThumbnail(c *gin.Context) { h.writeOpenAPIOutput(c, true) } + +func (h *Handler) writeOpenAPIOutput(c *gin.Context, thumbnail bool) { + generationID, outputID, ok := openAPIFileIDs(c, "outputID") + if !ok { + return + } + reader, object, err := h.service.OpenOutput(c.Request.Context(), currentAPIPrincipal(c).UserID, generationID, outputID, thumbnail) + if err != nil { + h.serviceError(c, err) + return + } + defer reader.Close() + c.Header("Content-Type", object.ContentType) + c.Header("Content-Disposition", "inline") + c.Status(http.StatusOK) + _, _ = io.Copy(c.Writer, reader) +} + +func openAPIFileIDs(c *gin.Context, resourceName string) (uint64, uint64, bool) { + if !queryNamesAllowed(c) { + writeError(c, http.StatusBadRequest, "invalid_request", "query parameters are invalid") + return 0, 0, false + } + generationID, ok := uintParam(c, "id") + if !ok { + return 0, 0, false + } + resourceID, ok := uintParam(c, resourceName) + return generationID, resourceID, ok +} + +func openAPIGenerationResponse(generation model.Generation) gin.H { + response := generationResponse(generation) + response["created_at"] = generation.CreatedAt.UTC() + if generation.CompletedAt != nil { + completedAt := generation.CompletedAt.UTC() + response["completed_at"] = completedAt + } + return response +} diff --git a/portal/handler/router.go b/portal/handler/router.go index b0e4bf5..be6b046 100644 --- a/portal/handler/router.go +++ b/portal/handler/router.go @@ -55,13 +55,13 @@ func NewRouter(sessions *session.Manager, authService *auth.Service, generationS c.Next() }) static.StaticFS("/", http.FS(staticFS)) - router.Use(handler.sessionMiddleware) - router.GET("/login", handler.loginPage) - router.GET("/", handler.appPage) - router.GET("/api-keys", handler.apiKeysPage) - router.GET("/generations/:id", handler.appPage) - router.GET("/ui/generations/:id/result", handler.resultFragment) - api := router.Group("/api") + sessionRoutes := router.Group("", handler.sessionMiddleware) + sessionRoutes.GET("/login", handler.loginPage) + sessionRoutes.GET("/", handler.appPage) + sessionRoutes.GET("/api-keys", handler.apiKeysPage) + sessionRoutes.GET("/generations/:id", handler.appPage) + sessionRoutes.GET("/ui/generations/:id/result", handler.resultFragment) + api := sessionRoutes.Group("/api") api.GET("/session", handler.sessionState) api.POST("/session/login", handler.csrf, handler.login) api.POST("/session/logout", handler.requireAuth, handler.csrf, handler.logout) @@ -82,6 +82,15 @@ func NewRouter(sessions *session.Manager, authService *auth.Service, generationS generations.GET("/:id/inputs/:inputID", handler.input) generations.GET("/:id/outputs/:outputID", handler.output) generations.GET("/:id/outputs/:outputID/thumbnail", handler.thumbnail) + openAPI := router.Group("/openapi/v1", handler.openAPIHeaders, handler.requireAPIKey) + openAPI.GET("/openapi.json", handler.openAPISpec) + openAPI.POST("/generations/text", handler.openAPISubmitText) + openAPI.POST("/generations/image", handler.openAPISubmitImages) + openAPI.GET("/generations", handler.openAPIHistory) + openAPI.GET("/generations/:id", handler.openAPIDetail) + openAPI.GET("/generations/:id/inputs/:inputID", handler.openAPIInput) + openAPI.GET("/generations/:id/outputs/:outputID", handler.openAPIOutput) + openAPI.GET("/generations/:id/outputs/:outputID/thumbnail", handler.openAPIThumbnail) return router, nil } @@ -388,6 +397,8 @@ func (h *Handler) serviceError(c *gin.Context, err error) { writeError(c, 400, "invalid_prompt", "prompt is invalid") case errors.Is(err, service.ErrInvalidIdempotencyKey): writeError(c, 400, "invalid_idempotency_key", "idempotency key is invalid") + case errors.Is(err, service.ErrIdempotencyConflict): + writeError(c, 409, "idempotency_conflict", "idempotency key was already used for a different request") case errors.Is(err, service.ErrImageCount): writeError(c, 400, "invalid_image_count", "image count is invalid") case errors.Is(err, service.ErrInvalidMetadata): diff --git a/portal/openapi/openapi.json b/portal/openapi/openapi.json new file mode 100644 index 0000000..1de77dd --- /dev/null +++ b/portal/openapi/openapi.json @@ -0,0 +1,200 @@ +{ + "openapi": "3.1.0", + "info": { + "title": "Chorus OpenAPI", + "version": "1.0.0", + "description": "Asynchronous text and image generation for existing controlled Chorus users. Submission never calls an upstream provider synchronously." + }, + "servers": [{"url": "/"}], + "security": [{"APIKeyAuth": []}], + "paths": { + "/openapi/v1/openapi.json": { + "get": { + "operationId": "getOpenAPISpecification", + "responses": { + "200": {"description": "OpenAPI 3.1 document", "content": {"application/json": {"schema": {"type": "object"}}}}, + "401": {"$ref": "#/components/responses/Unauthorized"} + } + } + }, + "/openapi/v1/generations/text": { + "post": { + "operationId": "createTextGeneration", + "parameters": [{"$ref": "#/components/parameters/IdempotencyKey"}], + "requestBody": { + "required": true, + "content": {"application/json": {"schema": {"$ref": "#/components/schemas/TextGenerationRequest"}}} + }, + "responses": { + "200": {"description": "Idempotent replay", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/Generation"}}}}, + "202": {"description": "Generation accepted", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/Generation"}}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"}, + "409": {"$ref": "#/components/responses/Conflict"}, + "503": {"$ref": "#/components/responses/Unavailable"} + } + } + }, + "/openapi/v1/generations/image": { + "post": { + "operationId": "createImageGeneration", + "description": "Multipart fields are prompt, metadata, and exactly one files[client_id] part for every image declared by metadata. image_generate has no file parts; image_edit requires at least one image and exactly one primary image.", + "parameters": [{"$ref": "#/components/parameters/IdempotencyKey"}], + "requestBody": { + "required": true, + "content": { + "multipart/form-data": { + "schema": { + "type": "object", + "required": ["prompt", "metadata"], + "properties": { + "prompt": {"type": "string"}, + "metadata": {"type": "string", "contentMediaType": "application/json", "contentSchema": {"$ref": "#/components/schemas/ImageMetadata"}} + }, + "additionalProperties": {"type": "string", "format": "binary"} + } + } + } + }, + "responses": { + "200": {"description": "Idempotent replay", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/Generation"}}}}, + "202": {"description": "Generation accepted", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/Generation"}}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"}, + "409": {"$ref": "#/components/responses/Conflict"}, + "503": {"$ref": "#/components/responses/Unavailable"} + } + } + }, + "/openapi/v1/generations": { + "get": { + "operationId": "listGenerations", + "parameters": [ + {"name": "limit", "in": "query", "schema": {"type": "integer", "minimum": 1, "maximum": 200}}, + {"name": "cursor", "in": "query", "schema": {"type": "string", "minLength": 1}} + ], + "responses": { + "200": {"description": "Generation history", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/GenerationPage"}}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"} + } + } + }, + "/openapi/v1/generations/{id}": { + "get": { + "operationId": "getGeneration", + "parameters": [{"$ref": "#/components/parameters/GenerationID"}], + "responses": { + "200": {"description": "Generation detail", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/GenerationDetail"}}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"}, + "404": {"$ref": "#/components/responses/NotFound"} + } + } + }, + "/openapi/v1/generations/{id}/inputs/{inputID}": { + "get": { + "operationId": "getGenerationInput", + "parameters": [{"$ref": "#/components/parameters/GenerationID"}, {"$ref": "#/components/parameters/InputID"}], + "responses": { + "200": {"description": "Original input file", "content": {"image/png": {}, "image/jpeg": {}, "image/webp": {}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"}, + "404": {"$ref": "#/components/responses/NotFound"} + } + } + }, + "/openapi/v1/generations/{id}/outputs/{outputID}": { + "get": { + "operationId": "getGenerationOutput", + "parameters": [{"$ref": "#/components/parameters/GenerationID"}, {"$ref": "#/components/parameters/OutputID"}], + "responses": { + "200": {"description": "Original output file", "content": {"image/png": {}, "image/jpeg": {}, "image/webp": {}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"}, + "404": {"$ref": "#/components/responses/NotFound"} + } + } + }, + "/openapi/v1/generations/{id}/outputs/{outputID}/thumbnail": { + "get": { + "operationId": "getGenerationOutputThumbnail", + "parameters": [{"$ref": "#/components/parameters/GenerationID"}, {"$ref": "#/components/parameters/OutputID"}], + "responses": { + "200": {"description": "Output thumbnail", "content": {"image/png": {}, "image/jpeg": {}, "image/webp": {}}}, + "400": {"$ref": "#/components/responses/BadRequest"}, + "401": {"$ref": "#/components/responses/Unauthorized"}, + "404": {"$ref": "#/components/responses/NotFound"} + } + } + } + }, + "components": { + "securitySchemes": { + "APIKeyAuth": {"type": "http", "scheme": "bearer", "bearerFormat": "Chorus API Key"} + }, + "parameters": { + "IdempotencyKey": {"name": "Idempotency-Key", "in": "header", "required": true, "schema": {"type": "string", "minLength": 8, "maxLength": 128, "pattern": "^[!-~](?:[ -~]{6,126}[!-~])?$"}}, + "GenerationID": {"name": "id", "in": "path", "required": true, "schema": {"type": "integer", "format": "int64", "minimum": 1}}, + "InputID": {"name": "inputID", "in": "path", "required": true, "schema": {"type": "integer", "format": "int64", "minimum": 1}}, + "OutputID": {"name": "outputID", "in": "path", "required": true, "schema": {"type": "integer", "format": "int64", "minimum": 1}} + }, + "responses": { + "BadRequest": {"description": "Invalid request", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorEnvelope"}}}}, + "Unauthorized": {"description": "Missing or invalid API key", "headers": {"WWW-Authenticate": {"schema": {"type": "string"}}, "X-Request-ID": {"schema": {"type": "string"}}}, "content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorEnvelope"}}}}, + "NotFound": {"description": "Resource not found", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorEnvelope"}}}}, + "Conflict": {"description": "Idempotency key conflict", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorEnvelope"}}}}, + "Unavailable": {"description": "Generation route unavailable", "content": {"application/json": {"schema": {"$ref": "#/components/schemas/ErrorEnvelope"}}}} + }, + "schemas": { + "TextGenerationRequest": {"type": "object", "additionalProperties": false, "required": ["prompt"], "properties": {"prompt": {"type": "string", "minLength": 1}}}, + "ImageMetadata": { + "type": "object", "additionalProperties": false, "required": ["capability", "role_rule", "images"], + "properties": { + "capability": {"type": "string", "enum": ["image_generate", "image_edit"]}, + "role_rule": {"type": "string"}, + "images": {"type": "array", "items": {"$ref": "#/components/schemas/ImageInput"}} + } + }, + "ImageInput": { + "type": "object", "additionalProperties": false, "required": ["client_id", "position", "role", "note"], + "properties": { + "client_id": {"type": "string", "pattern": "^[A-Za-z0-9_-]{1,64}$"}, + "position": {"type": "integer", "minimum": 0}, + "role": {"type": "string", "enum": ["primary", "reference"]}, + "note": {"type": "string", "maxLength": 500} + } + }, + "Generation": { + "type": "object", "required": ["id", "kind", "status", "terminal", "user_prompt", "created_at"], + "properties": { + "id": {"type": "integer", "format": "int64"}, + "kind": {"type": "string", "enum": ["text", "image"]}, + "status": {"type": "string", "enum": ["pending", "running", "succeeded", "failed"]}, + "terminal": {"type": "boolean"}, + "user_prompt": {"type": "string"}, + "error_code": {"type": ["string", "null"]}, + "error_message": {"type": ["string", "null"]}, + "created_at": {"type": "string", "format": "date-time"}, + "completed_at": {"type": ["string", "null"], "format": "date-time"} + } + }, + "GenerationDetail": { + "allOf": [ + {"$ref": "#/components/schemas/Generation"}, + {"type": "object", "required": ["inputs", "outputs"], "properties": {"role_rule": {"type": ["string", "null"]}, "inputs": {"type": "array", "items": {"$ref": "#/components/schemas/GenerationInput"}}, "outputs": {"type": "array", "items": {"$ref": "#/components/schemas/GenerationOutput"}}}} + ] + }, + "GenerationInput": { + "type": "object", + "properties": {"id": {"type": "integer", "format": "int64"}, "position": {"type": "integer"}, "role": {"type": "string"}, "note": {"type": ["string", "null"]}, "name": {"type": "string"}, "mime_type": {"type": "string"}, "size_bytes": {"type": "integer"}, "width": {"type": ["integer", "null"]}, "height": {"type": ["integer", "null"]}, "url": {"type": "string"}} + }, + "GenerationOutput": { + "type": "object", + "properties": {"id": {"type": "integer", "format": "int64"}, "kind": {"type": "string"}, "text": {"type": "string"}, "mime_type": {"type": ["string", "null"]}, "size_bytes": {"type": ["integer", "null"]}, "width": {"type": ["integer", "null"]}, "height": {"type": ["integer", "null"]}, "url": {"type": "string"}, "thumbnail_url": {"type": "string"}, "created_at": {"type": "string", "format": "date-time"}} + }, + "GenerationPage": {"type": "object", "required": ["items", "next_cursor", "has_more"], "properties": {"items": {"type": "array", "items": {"$ref": "#/components/schemas/Generation"}}, "next_cursor": {"type": "string"}, "has_more": {"type": "boolean"}}}, + "ErrorEnvelope": {"type": "object", "additionalProperties": false, "required": ["error"], "properties": {"error": {"type": "object", "additionalProperties": true, "required": ["code", "message"], "properties": {"code": {"type": "string"}, "message": {"type": "string"}}}}} + } + } +} diff --git a/portal/openapi/spec.go b/portal/openapi/spec.go new file mode 100644 index 0000000..ee9c912 --- /dev/null +++ b/portal/openapi/spec.go @@ -0,0 +1,10 @@ +package openapi + +import _ "embed" + +//go:embed openapi.json +var specification []byte + +func Spec() []byte { + return append([]byte(nil), specification...) +} diff --git a/portal/openapi/spec_test.go b/portal/openapi/spec_test.go new file mode 100644 index 0000000..521652e --- /dev/null +++ b/portal/openapi/spec_test.go @@ -0,0 +1,50 @@ +package openapi + +import ( + "encoding/json" + "testing" +) + +func TestSpecificationCoversRegisteredV1Operations(t *testing.T) { + var document struct { + OpenAPI string `json:"openapi"` + Security []map[string][]string `json:"security"` + Paths map[string]map[string]any `json:"paths"` + Components map[string]json.RawMessage `json:"components"` + } + if err := json.Unmarshal(Spec(), &document); err != nil { + t.Fatal(err) + } + if document.OpenAPI != "3.1.0" { + t.Fatalf("openapi version = %q", document.OpenAPI) + } + if len(document.Security) != 1 { + t.Fatalf("security = %#v", document.Security) + } + expected := map[string]string{ + "/openapi/v1/openapi.json": "get", + "/openapi/v1/generations/text": "post", + "/openapi/v1/generations/image": "post", + "/openapi/v1/generations": "get", + "/openapi/v1/generations/{id}": "get", + "/openapi/v1/generations/{id}/inputs/{inputID}": "get", + "/openapi/v1/generations/{id}/outputs/{outputID}": "get", + "/openapi/v1/generations/{id}/outputs/{outputID}/thumbnail": "get", + } + if len(document.Paths) != len(expected) { + t.Fatalf("path count = %d, want %d", len(document.Paths), len(expected)) + } + for path, method := range expected { + if _, exists := document.Paths[path][method]; !exists { + t.Errorf("missing %s %s", method, path) + } + } +} + +func TestSpecReturnsDefensiveCopy(t *testing.T) { + first := Spec() + first[0] = 'x' + if second := Spec(); len(second) == 0 || second[0] != '{' { + t.Fatal("Spec returned mutable embedded storage") + } +} diff --git a/portal/service/idempotency.go b/portal/service/idempotency.go new file mode 100644 index 0000000..b76a530 --- /dev/null +++ b/portal/service/idempotency.go @@ -0,0 +1,74 @@ +package service + +import ( + "bytes" + "context" + "encoding/json" + "io" + "strings" + + "git.ilapage.cn/OPC/chorus/internal/core/model" + corerouter "git.ilapage.cn/OPC/chorus/internal/core/router" +) + +func validIdempotencyKey(value string) bool { + if len(value) < 8 || len(value) > 128 || strings.TrimSpace(value) != value { + return false + } + for index := 0; index < len(value); index++ { + if value[index] < 0x20 || value[index] > 0x7e { + return false + } + } + return true +} + +func textReplayMatches(generation model.Generation, prompt string) bool { + return generation.Kind == model.KindText && generation.UserPrompt == prompt +} + +func (s *Service) imageReplayMatches(ctx context.Context, generation model.Generation, capability model.Capability, prompt, roleRule string, expected []validatedImage) (bool, error) { + if generation.Kind != model.KindImage || generation.UserPrompt != prompt || nullableString(generation.RoleRule) != roleRule { + return false, nil + } + var snapshot corerouter.RouteSnapshot + if json.Unmarshal(generation.RouteSnapshot, &snapshot) != nil || snapshot.Capability != capability { + return false, nil + } + var inputs []model.GenerationInput + if err := s.db.WithContext(ctx).Where("generation_id = ?", generation.ID).Order("position,id").Find(&inputs).Error; err != nil { + return false, err + } + if len(inputs) != len(expected) { + return false, nil + } + for index, input := range inputs { + want := expected[index] + if input.Position != want.Position || input.Role != want.Role || nullableString(input.Note) != want.Note || input.OriginalName != want.Name || input.MIMEType != want.MIME || input.SizeBytes != uint64(len(want.Content)) || input.Width == nil || input.Height == nil || *input.Width != uint32(want.Width) || *input.Height != uint32(want.Height) { + return false, nil + } + reader, object, err := s.storage.Open(ctx, input.StorageKey) + if err != nil { + return false, err + } + content, readErr := io.ReadAll(io.LimitReader(reader, s.config.MaxImageBytes+1)) + closeErr := reader.Close() + if readErr != nil { + return false, readErr + } + if closeErr != nil { + return false, closeErr + } + if object.OwnerID != generation.UserID || object.GenerationID != generation.ID || object.ContentType != want.MIME || !bytes.Equal(content, want.Content) { + return false, nil + } + } + return true, nil +} + +func nullableString(value *string) string { + if value == nil { + return "" + } + return *value +} diff --git a/portal/service/openapi.go b/portal/service/openapi.go new file mode 100644 index 0000000..d8d3a45 --- /dev/null +++ b/portal/service/openapi.go @@ -0,0 +1,65 @@ +package service + +import ( + "context" + "errors" + "fmt" + "time" + + coreapikey "git.ilapage.cn/OPC/chorus/internal/core/apikey" + "git.ilapage.cn/OPC/chorus/internal/core/model" + platformapikey "git.ilapage.cn/OPC/chorus/internal/platform/apikey" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +var ErrInvalidAPIKeyCredential = errors.New("invalid_api_key") + +type APIPrincipal struct { + UserID uint64 + APIKeyID uint64 +} + +func (s *Service) AuthenticateAPIKey(ctx context.Context, token string, now time.Time) (APIPrincipal, error) { + parsed, err := platformapikey.Parse(token) + if err != nil { + return APIPrincipal{}, ErrInvalidAPIKeyCredential + } + now = now.UTC() + var principal APIPrincipal + err = s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + repository, repositoryErr := coreapikey.NewGORMRepository(tx) + if repositoryErr != nil { + return repositoryErr + } + key, findErr := repository.ByPublicIDLocked(ctx, parsed.PublicID) + if errors.Is(findErr, coreapikey.ErrAPIKeyNotFound) || errors.Is(findErr, coreapikey.ErrInvalidAPIKey) { + return ErrInvalidAPIKeyCredential + } + if findErr != nil { + return findErr + } + if !platformapikey.Verify(token, key.PublicID, key.SecretHash) || !key.UsableAt(now) { + return ErrInvalidAPIKeyCredential + } + var user model.User + userErr := tx.Clauses(clause.Locking{Strength: "SHARE"}).Select("id", "status").First(&user, key.UserID).Error + if errors.Is(userErr, gorm.ErrRecordNotFound) || (userErr == nil && user.Status != "active") { + return ErrInvalidAPIKeyCredential + } + if userErr != nil { + return fmt.Errorf("verify API key user: %w", userErr) + } + principal = APIPrincipal{UserID: key.UserID, APIKeyID: key.ID} + return nil + }) + if err != nil { + return APIPrincipal{}, err + } + if updateErr := s.db.WithContext(ctx).Model(&model.APIKey{}). + Where("id = ? AND (last_used_at IS NULL OR last_used_at < ?)", principal.APIKeyID, now.Add(-time.Minute)). + Update("last_used_at", now).Error; updateErr != nil { + return APIPrincipal{}, fmt.Errorf("update API key last used time: %w", updateErr) + } + return principal, nil +} diff --git a/portal/service/service.go b/portal/service/service.go index 72f1df3..13abd2b 100644 --- a/portal/service/service.go +++ b/portal/service/service.go @@ -25,6 +25,7 @@ import ( var ( ErrInvalidPrompt = errors.New("invalid_prompt") ErrInvalidIdempotencyKey = errors.New("invalid_idempotency_key") + ErrIdempotencyConflict = errors.New("idempotency_conflict") ErrImageCount = errors.New("invalid_image_count") ErrImageTooLarge = errors.New("image_too_large") ErrUploadTooLarge = errors.New("upload_too_large") @@ -96,6 +97,9 @@ func (s *Service) SubmitText(ctx context.Context, userID uint64, idempotencyKey, return model.Generation{}, false, err } if existing, found, err := s.existingGeneration(ctx, userID, idempotencyKey); err != nil || found { + if err == nil && !textReplayMatches(existing, prompt) { + return model.Generation{}, false, ErrIdempotencyConflict + } return existing, false, err } snapshot, promptTemplate, err := s.submissionRoute(ctx, model.CapabilityText) @@ -111,6 +115,9 @@ func (s *Service) SubmitText(ctx context.Context, userID uint64, idempotencyKey, return model.Generation{}, false, err } created, err := s.creator.CreateIdempotentPrepared(ctx, &generation, func(uint64) ([]model.GenerationInput, error) { return nil, nil }) + if err == nil && !created && !textReplayMatches(generation, prompt) { + return model.Generation{}, false, ErrIdempotencyConflict + } return generation, created, err } @@ -118,13 +125,23 @@ func (s *Service) SubmitImages(ctx context.Context, userID uint64, idempotencyKe if err := s.validateCommon(userID, idempotencyKey, prompt); err != nil { return model.Generation{}, false, err } - 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 } + if existing, found, findErr := s.existingGeneration(ctx, userID, idempotencyKey); findErr != nil || found { + if findErr != nil { + return model.Generation{}, false, findErr + } + matches, matchErr := s.imageReplayMatches(ctx, existing, metadata.Capability, prompt, roleRule, validated) + if matchErr != nil { + return model.Generation{}, false, matchErr + } + if !matches { + return model.Generation{}, false, ErrIdempotencyConflict + } + return existing, false, nil + } snapshot, promptTemplate, err := s.submissionRoute(ctx, metadata.Capability) if err != nil { return model.Generation{}, false, err @@ -166,6 +183,15 @@ func (s *Service) SubmitImages(ctx context.Context, userID uint64, idempotencyKe _ = s.storage.Delete(context.Background(), key) } } + if err == nil && !created { + matches, matchErr := s.imageReplayMatches(ctx, generation, metadata.Capability, prompt, roleRule, validated) + if matchErr != nil { + return model.Generation{}, false, matchErr + } + if !matches { + return model.Generation{}, false, ErrIdempotencyConflict + } + } return generation, created, err } @@ -227,7 +253,7 @@ func (s *Service) validateCommon(userID uint64, key, prompt string) error { if userID == 0 { return ErrNotFound } - if len(key) < 8 || len(key) > 128 || strings.TrimSpace(key) != key { + if !validIdempotencyKey(key) { return ErrInvalidIdempotencyKey } if strings.TrimSpace(prompt) == "" || len(prompt) > s.config.MaxPromptBytes || !utf8.ValidString(prompt) { diff --git a/portal/service/validation_test.go b/portal/service/validation_test.go index c818d7d..2288eba 100644 --- a/portal/service/validation_test.go +++ b/portal/service/validation_test.go @@ -138,6 +138,42 @@ func TestWebPValidationMatchesConfirmedPrototype(t *testing.T) { } } +func TestIdempotencyKeyUsesPrintableASCII(t *testing.T) { + tests := []struct { + value string + want bool + }{ + {"12345678", true}, + {"request key 123", true}, + {strings.Repeat("x", 128), true}, + {"short", false}, + {" leading-key", false}, + {"trailing-key ", false}, + {"line\nbreak", false}, + {"中文幂等键值", false}, + {strings.Repeat("x", 129), false}, + } + for _, test := range tests { + if got := validIdempotencyKey(test.value); got != test.want { + t.Errorf("validIdempotencyKey(%q) = %v, want %v", test.value, got, test.want) + } + } +} + +func TestTextReplayMustMatchOriginalRequest(t *testing.T) { + generation := model.Generation{Kind: model.KindText, UserPrompt: "same prompt"} + if !textReplayMatches(generation, "same prompt") { + t.Fatal("matching text request was rejected") + } + if textReplayMatches(generation, "different prompt") { + t.Fatal("different prompt was accepted as an idempotent replay") + } + generation.Kind = model.KindImage + if textReplayMatches(generation, "same prompt") { + t.Fatal("different generation kind was accepted as an idempotent replay") + } +} + func smallPNG(width, height int) []byte { img := image.NewRGBA(image.Rect(0, 0, width, height)) for y := 0; y < height; y++ {