Files
chorus/portal/handler/router.go
T

544 lines
19 KiB
Go

package handler
import (
"encoding/json"
"errors"
"fmt"
"html/template"
"io"
"mime"
"net"
"net/http"
"net/url"
"path"
"strconv"
"strings"
"git.ilapage.cn/OPC/chorus/internal/core/model"
corerouter "git.ilapage.cn/OPC/chorus/internal/core/router"
"git.ilapage.cn/OPC/chorus/portal/auth"
"git.ilapage.cn/OPC/chorus/portal/service"
"git.ilapage.cn/OPC/chorus/portal/session"
"git.ilapage.cn/OPC/chorus/portal/web"
"github.com/gin-gonic/gin"
)
type Handler struct {
sessions *session.Manager
auth *auth.Service
service *service.Service
maxUploadBytes int64
renderer *web.Renderer
}
func NewRouter(sessions *session.Manager, authService *auth.Service, generationService *service.Service, maxUploadBytes int64) (*gin.Engine, error) {
if sessions == nil || authService == nil || generationService == nil || maxUploadBytes <= 0 {
return nil, errors.New("portal handler configuration is invalid")
}
handler := &Handler{sessions: sessions, auth: authService, service: generationService, maxUploadBytes: maxUploadBytes}
renderer, err := web.NewRenderer()
if err != nil {
return nil, err
}
handler.renderer = renderer
router := gin.New()
router.Use(gin.RecoveryWithWriter(io.Discard))
router.Use(securityHeaders)
_ = router.SetTrustedProxies(nil)
staticFS, err := web.Static()
if err != nil {
return nil, err
}
static := router.Group("/static")
static.Use(func(c *gin.Context) {
c.Header("Cache-Control", "public, max-age=3600")
c.Next()
})
static.StaticFS("/", http.FS(staticFS))
router.Use(handler.sessionMiddleware)
router.GET("/login", handler.loginPage)
router.GET("/", handler.appPage)
router.GET("/generations/:id", handler.appPage)
router.GET("/ui/generations/:id/result", handler.resultFragment)
api := router.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)
generations := api.Group("/generations", handler.requireAuth)
generations.GET("", handler.history)
generations.POST("/text", handler.csrf, handler.submitText)
generations.POST("/image", handler.csrf, handler.submitImages)
generations.GET("/:id", handler.detail)
generations.GET("/:id/card", handler.card)
generations.GET("/:id/inputs/:inputID", handler.input)
generations.GET("/:id/outputs/:outputID", handler.output)
generations.GET("/:id/outputs/:outputID/thumbnail", handler.thumbnail)
return router, nil
}
func (h *Handler) sessionMiddleware(c *gin.Context) {
state, err := h.sessions.Ensure(c.Writer, c.Request)
if err != nil {
writeError(c, http.StatusInternalServerError, "internal_error", "request could not be completed")
c.Abort()
return
}
c.Set("session", state)
c.Next()
}
func (h *Handler) requireAuth(c *gin.Context) {
state := currentSession(c)
if state.UserID == 0 {
if isHTMX(c) {
c.Header("HX-Redirect", "/login")
}
writeError(c, http.StatusUnauthorized, "authentication_required", "sign in is required")
c.Abort()
return
}
c.Next()
}
func (h *Handler) csrf(c *gin.Context) {
if !h.sessions.ValidateCSRF(c.Request, currentSession(c)) {
writeError(c, http.StatusForbidden, "csrf_invalid", "request token is invalid")
c.Abort()
return
}
c.Next()
}
func currentSession(c *gin.Context) session.State {
value, _ := c.Get("session")
state, _ := value.(session.State)
return state
}
func (h *Handler) sessionState(c *gin.Context) {
state := currentSession(c)
c.JSON(http.StatusOK, gin.H{"authenticated": state.UserID != 0, "csrf_token": state.CSRFToken})
}
func (h *Handler) login(c *gin.Context) {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 64<<10)
var input struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required"`
}
if err := c.ShouldBindJSON(&input); err != nil {
writeError(c, 400, "invalid_request", "email and password are required")
return
}
user, err := h.auth.Login(c.Request.Context(), remoteKey(c.Request), input.Email, input.Password)
if err != nil {
if !errors.Is(err, auth.ErrInvalidCredentials) {
writeError(c, http.StatusInternalServerError, "internal_error", "request could not be completed")
return
}
writeError(c, http.StatusUnauthorized, "invalid_credentials", "email or password is incorrect")
return
}
state, err := h.sessions.Authenticate(c.Writer, c.Request, user.ID)
if err != nil {
writeError(c, 500, "internal_error", "request could not be completed")
return
}
c.JSON(http.StatusOK, gin.H{"authenticated": true, "csrf_token": state.CSRFToken, "user": gin.H{"id": user.ID, "display_name": user.DisplayName}})
}
func (h *Handler) logout(c *gin.Context) {
state, err := h.sessions.Logout(c.Writer, c.Request)
if err != nil {
writeError(c, 500, "internal_error", "request could not be completed")
return
}
c.JSON(http.StatusOK, gin.H{"authenticated": false, "csrf_token": state.CSRFToken})
}
func (h *Handler) submitText(c *gin.Context) {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, 64<<10)
var input struct {
IdempotencyKey string `json:"idempotency_key"`
Prompt string `json:"prompt"`
}
if c.ShouldBindJSON(&input) != nil {
writeError(c, 400, "invalid_request", "request body is invalid")
return
}
generation, created, err := h.service.SubmitText(c.Request.Context(), currentSession(c).UserID, input.IdempotencyKey, input.Prompt)
if err != nil {
h.serviceError(c, err)
return
}
status := http.StatusAccepted
if !created {
status = http.StatusOK
}
c.JSON(status, generationResponse(generation))
}
func (h *Handler) submitImages(c *gin.Context) {
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, 400, "upload_too_large", "upload exceeds configured limit")
return
}
defer c.Request.MultipartForm.RemoveAll()
var metadata service.ImageMetadata
decoder := json.NewDecoder(strings.NewReader(c.PostForm("metadata")))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&metadata); err != nil || decoder.Decode(&struct{}{}) != io.EOF {
writeError(c, 400, "invalid_metadata", "image metadata is invalid")
return
}
expectedFields := make(map[string]string, len(metadata.Images))
for _, item := range metadata.Images {
field := "files[" + item.ClientID + "]"
if _, exists := expectedFields[field]; exists {
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
return
}
expectedFields[field] = item.ClientID
}
for field := range c.Request.MultipartForm.File {
if _, ok := expectedFields[field]; !ok {
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
return
}
}
uploads := make([]service.Upload, 0, len(expectedFields))
for field, clientID := range expectedFields {
headers := c.Request.MultipartForm.File[field]
if len(headers) != 1 {
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
return
}
header := headers[0]
file, err := header.Open()
if err != nil {
writeFileError(c, int(metadataPosition(metadata, clientID)), safeName(header.Filename), "invalid_image_content")
return
}
content, readErr := io.ReadAll(io.LimitReader(file, h.maxUploadBytes+1))
file.Close()
if readErr != nil {
writeFileError(c, int(metadataPosition(metadata, clientID)), safeName(header.Filename), "invalid_image_content")
return
}
uploads = append(uploads, service.Upload{ClientID: clientID, Name: header.Filename, DeclaredMIME: header.Header.Get("Content-Type"), Content: content})
}
generation, created, err := h.service.SubmitImages(c.Request.Context(), currentSession(c).UserID, c.PostForm("idempotency_key"), c.PostForm("prompt"), metadata, uploads)
if err != nil {
h.serviceError(c, err)
return
}
status := http.StatusAccepted
if !created {
status = http.StatusOK
}
c.JSON(status, generationResponse(generation))
}
func metadataPosition(metadata service.ImageMetadata, clientID string) uint32 {
for _, item := range metadata.Images {
if item.ClientID == clientID {
return item.Position
}
}
return 0
}
func (h *Handler) history(c *gin.Context) {
limit := h.service.DefaultHistoryLimit()
if raw, exists := c.GetQuery("limit"); exists {
parsed, parseErr := strconv.Atoi(raw)
if parseErr != nil {
writeError(c, 400, "invalid_limit", "history limit is invalid")
return
}
limit = parsed
}
cursor, hasCursor := c.GetQuery("cursor")
if hasCursor && cursor == "" {
writeError(c, 400, "invalid_cursor", "history cursor is invalid")
return
}
page, err := h.service.History(c.Request.Context(), currentSession(c).UserID, limit, cursor)
if err != nil {
h.serviceError(c, err)
return
}
items := make([]gin.H, 0, len(page.Items))
for _, row := range page.Items {
items = append(items, generationResponse(row))
}
c.JSON(200, gin.H{"items": items, "next_cursor": page.NextCursor, "has_more": page.HasMore})
}
func (h *Handler) detail(c *gin.Context) {
id, ok := uintParam(c, "id")
if !ok {
return
}
detail, err := h.service.ByID(c.Request.Context(), currentSession(c).UserID, id)
if err != nil {
h.serviceError(c, err)
return
}
outputs := make([]gin.H, 0, len(detail.Outputs))
for _, output := range detail.Outputs {
item := gin.H{"id": output.ID, "kind": output.Kind}
if output.TextContent != nil {
item["text"] = *output.TextContent
}
if output.StorageKey != nil {
item["url"] = fmt.Sprintf("/api/generations/%d/outputs/%d", id, output.ID)
item["thumbnail_url"] = fmt.Sprintf("/api/generations/%d/outputs/%d/thumbnail", id, output.ID)
}
outputs = append(outputs, item)
}
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, "url": fmt.Sprintf("/api/generations/%d/inputs/%d", id, input.ID)})
}
response := generationResponse(detail.Generation)
response["role_rule"] = detail.Generation.RoleRule
response["inputs"] = inputs
response["outputs"] = outputs
c.JSON(200, response)
}
func (h *Handler) card(c *gin.Context) {
id, ok := uintParam(c, "id")
if !ok {
return
}
detail, err := h.service.ByID(c.Request.Context(), currentSession(c).UserID, id)
if err != nil {
h.serviceError(c, err)
return
}
data := struct {
ID uint64
Status model.GenerationStatus
Terminal bool
}{id, detail.Generation.Status, detail.Generation.Status.Terminal()}
const markup = `<article data-generation-id="{{.ID}}" data-status="{{.Status}}" data-terminal="{{.Terminal}}"><span>{{.Status}}</span></article>`
tmpl := template.Must(template.New("card").Parse(markup))
c.Header("Content-Type", "text/html; charset=utf-8")
_ = tmpl.Execute(c.Writer, data)
}
func (h *Handler) input(c *gin.Context) {
generationID, ok := uintParam(c, "id")
if !ok {
return
}
inputID, ok := uintParam(c, "inputID")
if !ok {
return
}
reader, object, name, err := h.service.OpenInput(c.Request.Context(), currentSession(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(200)
_, _ = io.Copy(c.Writer, reader)
}
func (h *Handler) output(c *gin.Context) { h.writeOutput(c, false) }
func (h *Handler) thumbnail(c *gin.Context) { h.writeOutput(c, true) }
func (h *Handler) writeOutput(c *gin.Context, thumbnail bool) {
generationID, ok := uintParam(c, "id")
if !ok {
return
}
outputID, ok := uintParam(c, "outputID")
if !ok {
return
}
reader, object, err := h.service.OpenOutput(c.Request.Context(), currentSession(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(200)
_, _ = io.Copy(c.Writer, reader)
}
func generationResponse(generation model.Generation) gin.H {
return gin.H{"id": generation.ID, "kind": generation.Kind, "status": generation.Status, "terminal": generation.Status.Terminal(), "user_prompt": generation.UserPrompt, "error_code": generation.ErrorCode, "error_message": generation.ErrorMessage, "created_at": generation.CreatedAt, "completed_at": generation.CompletedAt}
}
func (h *Handler) serviceError(c *gin.Context, err error) {
var fileError *service.FileError
if errors.As(err, &fileError) {
writeFileError(c, fileError.Index, fileError.Name, fileError.Code.Error())
return
}
switch {
case errors.Is(err, service.ErrInvalidPrompt):
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.ErrImageCount):
writeError(c, 400, "invalid_image_count", "image count is invalid")
case errors.Is(err, service.ErrInvalidMetadata):
writeError(c, 400, "invalid_metadata", "image metadata is invalid")
case errors.Is(err, service.ErrInvalidCapability):
writeError(c, 400, "invalid_capability", "generation capability is invalid")
case errors.Is(err, service.ErrInvalidClientID):
writeError(c, 400, "invalid_client_id", "image client id is invalid")
case errors.Is(err, service.ErrImageMapping):
writeError(c, 400, "invalid_image_mapping", "image files do not match metadata")
case errors.Is(err, service.ErrImagePosition):
writeError(c, 400, "invalid_image_position", "image positions must be continuous")
case errors.Is(err, service.ErrImageRole):
writeError(c, 400, "invalid_image_role", "image role is invalid")
case errors.Is(err, service.ErrImagePrimary):
writeError(c, 400, "invalid_primary_count", "image edit requires exactly one primary image")
case errors.Is(err, service.ErrInvalidRoleRule):
writeError(c, 400, "invalid_role_rule", "image role rule is invalid")
case errors.Is(err, service.ErrInvalidNote):
writeError(c, 400, "invalid_image_note", "image note is invalid")
case errors.Is(err, service.ErrInvalidLimit):
writeError(c, 400, "invalid_limit", "history limit is invalid")
case errors.Is(err, service.ErrInvalidCursor):
writeError(c, 400, "invalid_cursor", "history cursor is invalid")
case errors.Is(err, service.ErrCursorExpired):
writeError(c, 400, "cursor_expired", "history cursor has expired")
case errors.Is(err, corerouter.ErrRouteNotConfigured):
writeError(c, 503, "route_not_configured", "generation route is not configured")
case errors.Is(err, corerouter.ErrRouteUnavailable):
writeError(c, 503, "route_unavailable", "generation route is unavailable")
case errors.Is(err, service.ErrNotFound):
writeError(c, 404, "not_found", "resource was not found")
default:
writeError(c, 500, "internal_error", "request could not be completed")
}
}
func writeFileError(c *gin.Context, index int, name, code string) {
c.JSON(400, gin.H{"error": gin.H{"code": code, "message": "image is invalid", "file": gin.H{"index": index, "name": name}}})
}
func writeError(c *gin.Context, status int, code, message string) {
c.JSON(status, gin.H{"error": gin.H{"code": code, "message": message}})
}
func uintParam(c *gin.Context, name string) (uint64, bool) {
value, err := strconv.ParseUint(c.Param(name), 10, 64)
if err != nil || value == 0 {
writeError(c, 400, "invalid_id", "resource id is invalid")
return 0, false
}
return value, true
}
func remoteKey(request *http.Request) string {
host, _, err := net.SplitHostPort(request.RemoteAddr)
if err == nil {
return host
}
return strings.TrimSpace(request.RemoteAddr)
}
func isHTMX(c *gin.Context) bool { return strings.EqualFold(c.GetHeader("HX-Request"), "true") }
func safeName(value string) string {
name := path.Base(strings.ReplaceAll(strings.TrimSpace(value), "\\", "/"))
if name == "." || name == "" {
return "image"
}
return name
}
func securityHeaders(c *gin.Context) {
c.Header("Content-Security-Policy", "default-src 'self'; script-src 'self'; style-src 'self'; img-src 'self' blob:; connect-src 'self'; base-uri 'none'; frame-ancestors 'none'; form-action 'self'")
c.Header("Referrer-Policy", "no-referrer")
c.Header("X-Content-Type-Options", "nosniff")
c.Header("X-Frame-Options", "DENY")
c.Next()
}
func (h *Handler) loginPage(c *gin.Context) {
state := currentSession(c)
returnTo := safeReturnTo(c.Query("return_to"))
if state.UserID != 0 {
c.Redirect(http.StatusSeeOther, returnTo)
return
}
c.Header("Cache-Control", "no-store")
c.Header("Content-Type", "text/html; charset=utf-8")
if err := h.renderer.Render(c.Writer, "login", web.Page{Title: "登录", CSRFToken: state.CSRFToken, ReturnTo: returnTo}); err != nil {
c.Status(http.StatusInternalServerError)
}
}
func (h *Handler) appPage(c *gin.Context) {
state := currentSession(c)
if state.UserID == 0 {
c.Redirect(http.StatusSeeOther, "/login?return_to="+url.QueryEscape(c.Request.URL.RequestURI()))
return
}
user, err := h.service.User(c.Request.Context(), state.UserID)
if err != nil {
c.Status(http.StatusUnauthorized)
return
}
history, err := h.service.RecentHistory(c.Request.Context(), state.UserID)
if err != nil {
c.Status(http.StatusInternalServerError)
return
}
var current *service.Detail
if rawID := c.Param("id"); rawID != "" {
id, parseErr := strconv.ParseUint(rawID, 10, 64)
if parseErr != nil || id == 0 {
c.Status(http.StatusNotFound)
return
}
detail, detailErr := h.service.ByID(c.Request.Context(), state.UserID, id)
if detailErr != nil {
c.Status(http.StatusNotFound)
return
}
current = &detail
}
maxPromptBytes, maxImages := h.service.Limits()
page := web.Page{Title: "生成工作台", DisplayName: user.DisplayName, CSRFToken: state.CSRFToken, History: history, Current: current, MaxPromptBytes: maxPromptBytes, MaxImages: maxImages}
c.Header("Cache-Control", "no-store")
c.Header("Content-Type", "text/html; charset=utf-8")
if err := h.renderer.Render(c.Writer, "app", page); err != nil {
c.Status(http.StatusInternalServerError)
}
}
func (h *Handler) resultFragment(c *gin.Context) {
state := currentSession(c)
if state.UserID == 0 {
c.Header("HX-Trigger", `{"authExpired":true}`)
c.Status(http.StatusUnauthorized)
return
}
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil || id == 0 {
c.Status(http.StatusNotFound)
return
}
detail, err := h.service.ByID(c.Request.Context(), state.UserID, id)
if err != nil {
c.Status(http.StatusNotFound)
return
}
c.Header("Cache-Control", "no-store")
c.Header("Content-Type", "text/html; charset=utf-8")
if err := h.renderer.Render(c.Writer, "result", web.Page{Current: &detail}); err != nil {
c.Status(http.StatusInternalServerError)
}
}
func safeReturnTo(value string) string {
if value == "" || !strings.HasPrefix(value, "/") || strings.HasPrefix(value, "//") || strings.ContainsAny(value, "\\\r\n\t") {
return "/"
}
parsed, err := url.ParseRequestURI(value)
if err != nil || parsed.IsAbs() || parsed.Host != "" {
return "/"
}
return value
}