Files
chorus/portal/handler/openapi.go
T

413 lines
13 KiB
Go

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) {
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.Set(auditGenerationIDContextKey, generation.ID)
c.Set(auditGenerationCreatedContextKey, created)
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.Set(auditGenerationIDContextKey, generation.ID)
c.Set(auditGenerationCreatedContextKey, created)
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
}