413 lines
13 KiB
Go
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
|
|
}
|