diff --git a/internal/handler/message.go b/internal/handler/message.go index c7a34c9a2..c2d69d6ac 100644 --- a/internal/handler/message.go +++ b/internal/handler/message.go @@ -12,6 +12,7 @@ import ( "github.com/Tencent/WeKnora/internal/errors" "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/storageurl" "github.com/Tencent/WeKnora/internal/types" "github.com/Tencent/WeKnora/internal/types/interfaces" secutils "github.com/Tencent/WeKnora/internal/utils" @@ -21,19 +22,45 @@ import ( // It provides endpoints for loading and managing message history type MessageHandler struct { MessageService interfaces.MessageService // Service that implements message business logic + // FileService and StorageResolver back the optional `resource_urls=public` + // mode, which returns loadable HTTP URLs instead of internal + // `resource://` handles. Both may be nil, in which case only the default + // handle mode is available. + FileService interfaces.FileService + StorageResolver interfaces.StorageBackendResolver } // NewMessageHandler creates a new message handler instance with the required service // Parameters: // - messageService: Service that implements message business logic +// - fileService: Storage access used to sign public resource URLs +// - storageResolver: Resolves per-tenant storage backends for those URLs // // Returns a pointer to a new MessageHandler -func NewMessageHandler(messageService interfaces.MessageService) *MessageHandler { +func NewMessageHandler( + messageService interfaces.MessageService, + fileService interfaces.FileService, + storageResolver interfaces.StorageBackendResolver, +) *MessageHandler { return &MessageHandler{ - MessageService: messageService, + MessageService: messageService, + FileService: fileService, + StorageResolver: storageResolver, } } +// resolveResourceRewriter builds the storage-reference rewriter for one response +// from the request's `resource_urls` parameter, falling back to the deployment +// default. An invalid value is a client error the caller must turn into a 400. +func (h *MessageHandler) resolveResourceRewriter(c *gin.Context) (*storageurl.Rewriter, error) { + ctx := c.Request.Context() + mode, err := storageurl.ResolveMode(ctx, c.Query(storageurl.QueryParam)) + if err != nil { + return nil, err + } + return storageurl.NewRequestRewriter(ctx, mode, h.FileService, h.StorageResolver), nil +} + // LoadMessages godoc // @Summary 加载消息历史 // @Description 加载会话的消息历史,支持分页和时间筛选 @@ -61,9 +88,16 @@ func (h *MessageHandler) LoadMessages(c *gin.Context) { logger.Infof(ctx, "Loading messages params, session ID: %s, limit: %s, before time: %s", sessionID, limit, beforeTimeStr) - // Parse limit parameter with fallback to default - limitInt, err := strconv.Atoi(limit) + rewriter, err := h.resolveResourceRewriter(c) if err != nil { + logger.Warnf(ctx, "Invalid resource URL mode: %v", err) + c.Error(errors.NewBadRequestError(err.Error())) + return + } + + // Parse limit parameter with fallback to default + limitInt, convErr := strconv.Atoi(limit) + if convErr != nil { logger.Warnf(ctx, "Invalid limit value, using default value 20, input: %s", limit) limitInt = 20 } @@ -92,6 +126,7 @@ func (h *MessageHandler) LoadMessages(c *gin.Context) { "Successfully retrieved recent messages, session ID: %s, message count: %d", sessionID, len(messages), ) + rewriter.RewriteMessages(ctx, messages) c.JSON(http.StatusOK, gin.H{ "success": true, "data": messages, @@ -132,6 +167,7 @@ func (h *MessageHandler) LoadMessages(c *gin.Context) { "Successfully retrieved messages before time, session ID: %s, message count: %d", sessionID, len(messages), ) + rewriter.RewriteMessages(ctx, messages) c.JSON(http.StatusOK, gin.H{ "success": true, "data": messages, diff --git a/internal/handler/session/qa.go b/internal/handler/session/qa.go index d795cae98..b4d1e5a75 100644 --- a/internal/handler/session/qa.go +++ b/internal/handler/session/qa.go @@ -16,6 +16,7 @@ import ( "github.com/Tencent/WeKnora/internal/errors" "github.com/Tencent/WeKnora/internal/event" "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/storageurl" "github.com/Tencent/WeKnora/internal/types" secutils "github.com/Tencent/WeKnora/internal/utils" "github.com/gin-gonic/gin" @@ -51,6 +52,10 @@ type qaRequestContext struct { attachmentIDs []string // Pre-uploaded session-scoped document IDs, resolved after SSE starts attachmentMetas types.MessageAttachments // Metadata-only view of attachmentIDs for the persisted user message suggestionAttribution *types.SuggestionAttribution + // resourceRewriter turns internal storage references in the outbound stream + // into directly loadable URLs when the caller asks for `resource_urls=public`. + // Disabled (a pass-through) in the default handle mode. + resourceRewriter *storageurl.StreamRewriter // Snapshot of the request fields needed to persist the input-bar state // for session restoration. Kept verbatim from the request so we record @@ -109,6 +114,14 @@ func (h *Handler) parseQARequest(c *gin.Context, logPrefix string) (*qaRequestCo logger.Error(ctx, "Query content is empty") return nil, nil, errors.NewBadRequestError("Query content cannot be empty") } + + // Resolve the storage-reference representation up front: once the SSE stream + // has started an invalid value can no longer be reported as a 400. + resourceRewriter, err := h.resolveStreamRewriter(c) + if err != nil { + logger.Warnf(ctx, "Invalid resource URL mode: %v", err) + return nil, nil, errors.NewBadRequestError(err.Error()) + } if h.suggestionService != nil && request.SuggestionAttribution != nil { if err := h.suggestionService.ValidateAttribution(ctx, sessionID, request.Query, request.SuggestionAttribution); err != nil { return nil, nil, errors.NewBadRequestError("invalid suggestion attribution") @@ -360,6 +373,7 @@ func (h *Handler) parseQARequest(c *gin.Context, logPrefix string) (*qaRequestCo suggestionAttribution: request.SuggestionAttribution, reqAgentEnabled: request.AgentEnabled, reqAgentID: request.AgentID, + resourceRewriter: resourceRewriter, } return reqCtx, &request, nil @@ -702,9 +716,15 @@ func (h *Handler) SearchKnowledge(c *gin.Context) { } logger.Infof(ctx, "Knowledge search completed, found %d results", len(searchResults)) + rewriter, err := h.resolveResourceRewriter(c) + if err != nil { + logger.Warnf(ctx, "Invalid resource URL mode: %v", err) + c.Error(errors.NewBadRequestError(err.Error())) + return + } c.JSON(http.StatusOK, gin.H{ "success": true, - "data": searchResults, + "data": rewriter.CopyReferences(ctx, searchResults), }) } @@ -980,7 +1000,7 @@ func (h *Handler) executeQA(reqCtx *qaRequestContext, mode qaMode, generateTitle // Handle SSE events (blocking) shouldWaitForTitle := generateTitle && reqCtx.session.Title == "" h.handleAgentEventsForSSE(ctx, reqCtx.c, sessionID, reqCtx.assistantMessage.ID, - reqCtx.requestID, streamCtx.eventBus, shouldWaitForTitle) + reqCtx.requestID, streamCtx.eventBus, shouldWaitForTitle, reqCtx.resourceRewriter) } // runVLMAnalysisIfNeeded runs VLM image analysis within the async goroutine, diff --git a/internal/handler/session/resource_urls.go b/internal/handler/session/resource_urls.go new file mode 100644 index 000000000..73315a2b5 --- /dev/null +++ b/internal/handler/session/resource_urls.go @@ -0,0 +1,132 @@ +package session + +import ( + "context" + "strings" + + "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/storageurl" + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" + "github.com/gin-gonic/gin" +) + +// resolveResourceRewriter builds the storage-reference rewriter for one response +// from the request's `resource_urls` parameter, falling back to the deployment +// default. An invalid value is a client error the caller must turn into a 400. +// +// In the default handle mode the returned rewriter is disabled, so responses are +// left exactly as before. +func (h *Handler) resolveResourceRewriter(c *gin.Context) (*storageurl.Rewriter, error) { + ctx := c.Request.Context() + mode, err := storageurl.ResolveMode(ctx, c.Query(storageurl.QueryParam)) + if err != nil { + return nil, err + } + return storageurl.NewRequestRewriter(ctx, mode, h.fileService, h.storageResolver), nil +} + +// resolveStreamRewriter is resolveResourceRewriter plus the holdback buffer an +// SSE response needs, because a storage reference can straddle two deltas. It +// must be called before any SSE header is written so an invalid value is still +// reportable as a normal 400. +func (h *Handler) resolveStreamRewriter(c *gin.Context) (*storageurl.StreamRewriter, error) { + rewriter, err := h.resolveResourceRewriter(c) + if err != nil { + return nil, err + } + return storageurl.NewStreamRewriter(rewriter), nil +} + +// deltaResponseTypes are the SSE events whose Content is an incremental chunk +// that clients accumulate. A storage reference can straddle two chunks, so these +// go through the holdback buffer; every other event carries a complete value. +var deltaResponseTypes = map[types.ResponseType]bool{ + types.ResponseTypeAnswer: true, + types.ResponseTypeThinking: true, + types.ResponseTypeReflection: true, +} + +// holdbackKey identifies one delta stream. The event id is the key clients +// accumulate on, so interleaved answer and thinking streams hold back +// independently; the type prefix lets a flushed remainder be re-emitted as the +// event type it came from. +func holdbackKey(responseType types.ResponseType, eventID string) string { + return string(responseType) + "\x00" + eventID +} + +func parseHoldbackKey(key string) (types.ResponseType, string) { + responseType, eventID, _ := strings.Cut(key, "\x00") + return types.ResponseType(responseType), eventID +} + +// buildStreamResponseFor builds the SSE payload for evt and, in public mode, +// replaces storage references with URLs the client can load directly. +func buildStreamResponseFor( + ctx context.Context, + evt interfaces.StreamEvent, + requestID string, + rewriter *storageurl.StreamRewriter, +) *types.StreamResponse { + response := buildStreamResponse(evt, requestID) + if !rewriter.Enabled() { + return response + } + + if deltaResponseTypes[evt.Type] { + response.Content = rewriter.Push(ctx, holdbackKey(evt.Type, evt.ID), response.Content, evt.Done) + } else { + response.Content = rewriter.Rewriter().String(ctx, response.Content) + } + response.KnowledgeReferences = rewriter.Rewriter().CopyReferences(ctx, response.KnowledgeReferences) + response.Data = rewriter.Rewriter().CopyData(ctx, response.Data) + return response +} + +// emitStreamEvent writes one SSE payload. Content still sitting in the holdback +// buffer is released first when evt terminates the stream, because clients treat +// the completion marker as the end of the message. +func emitStreamEvent( + ctx context.Context, + c *gin.Context, + evt interfaces.StreamEvent, + requestID string, + rewriter *storageurl.StreamRewriter, +) { + response := buildStreamResponseFor(ctx, evt, requestID, rewriter) + if evt.Type == types.ResponseTypeComplete { + flushHeldStreamContent(ctx, c, requestID, rewriter) + } + c.SSEvent("message", response) + c.Writer.Flush() +} + +// flushHeldStreamContent emits whatever the holdback buffer still retains, so a +// trailing reference is not dropped when a delta stream ends without a terminal +// chunk. Callers invoke it once the stream is over; if the client has already +// gone there is nobody left to receive it. +func flushHeldStreamContent( + ctx context.Context, + c *gin.Context, + requestID string, + rewriter *storageurl.StreamRewriter, +) { + held := rewriter.FlushAll(ctx) + if len(held) == 0 || c.Request.Context().Err() != nil { + return + } + for key, content := range held { + if content == "" { + continue + } + responseType, eventID := parseHoldbackKey(key) + logger.Debugf(ctx, "Flushing held stream fragment, type: %s, event: %s", responseType, eventID) + c.SSEvent("message", &types.StreamResponse{ + ID: requestID, + ResponseType: responseType, + Content: content, + Data: map[string]interface{}{"event_id": eventID}, + }) + c.Writer.Flush() + } +} diff --git a/internal/handler/session/stream.go b/internal/handler/session/stream.go index cf7e8f3d7..d46cc42fd 100644 --- a/internal/handler/session/stream.go +++ b/internal/handler/session/stream.go @@ -10,6 +10,7 @@ import ( "github.com/Tencent/WeKnora/internal/errors" "github.com/Tencent/WeKnora/internal/event" "github.com/Tencent/WeKnora/internal/logger" + "github.com/Tencent/WeKnora/internal/storageurl" "github.com/Tencent/WeKnora/internal/types" "github.com/Tencent/WeKnora/internal/types/interfaces" secutils "github.com/Tencent/WeKnora/internal/utils" @@ -53,9 +54,17 @@ func (h *Handler) ContinueStream(c *gin.Context) { logger.Infof(ctx, "Continuing stream, session ID: %s, message ID: %s", sessionID, messageID) - // Verify that the session exists and belongs to this tenant - _, err := h.sessionService.GetSession(ctx, sessionID) + // Resolve before any SSE header is written so an invalid resource_urls value + // is still reportable as a normal 400 JSON error. + resourceRewriter, err := h.resolveStreamRewriter(c) if err != nil { + logger.Warnf(ctx, "Invalid resource URL mode: %v", err) + c.Error(errors.NewBadRequestError(err.Error())) + return + } + + // Verify that the session exists and belongs to this tenant + if _, err := h.sessionService.GetSession(ctx, sessionID); err != nil { if stderrors.Is(err, errors.ErrSessionNotFound) { logger.Warnf(ctx, "Session not found, ID: %s", sessionID) c.Error(errors.NewNotFoundError(err.Error())) @@ -139,9 +148,7 @@ func (h *Handler) ContinueStream(c *gin.Context) { // Replay existing events logger.Debugf(ctx, "Replaying %d existing events", len(events)) for _, evt := range events { - response := buildStreamResponse(evt, message.RequestID) - c.SSEvent("message", response) - c.Writer.Flush() + emitStreamEvent(ctx, c, evt, message.RequestID, resourceRewriter) } // If stream is already completed, send final event and return @@ -178,9 +185,7 @@ func (h *Handler) ContinueStream(c *gin.Context) { streamCompletedNow = true } - response := buildStreamResponse(evt, message.RequestID) - c.SSEvent("message", response) - c.Writer.Flush() + emitStreamEvent(ctx, c, evt, message.RequestID, resourceRewriter) } // Update offset @@ -326,6 +331,7 @@ func (h *Handler) handleAgentEventsForSSE( sessionID, assistantMessageID, requestID string, eventBus *event.EventBus, waitForTitle bool, + resourceRewriter *storageurl.StreamRewriter, ) { ticker := time.NewTicker(100 * time.Millisecond) defer ticker.Stop() @@ -387,7 +393,7 @@ func (h *Handler) handleAgentEventsForSSE( } // Build StreamResponse from StreamEvent - response := buildStreamResponse(evt, requestID) + response := buildStreamResponseFor(ctx, evt, requestID, resourceRewriter) // Check for completion event if evt.Type == "complete" { @@ -405,6 +411,12 @@ func (h *Handler) handleAgentEventsForSSE( return } + // Any content still held back must precede the completion marker, + // which clients treat as the end of the message. + if streamCompleted { + flushHeldStreamContent(ctx, c, requestID, resourceRewriter) + } + c.SSEvent("message", response) c.Writer.Flush() } @@ -436,9 +448,7 @@ func (h *Handler) handleAgentEventsForSSE( } if len(events) > 0 { for _, evt := range events { - response := buildStreamResponse(evt, requestID) - c.SSEvent("message", response) - c.Writer.Flush() + emitStreamEvent(ctx, c, evt, requestID, resourceRewriter) // If we got the title, we can exit if evt.Type == types.ResponseTypeSessionTitle { log.Infof("Title event received: %s", evt.Content) diff --git a/internal/storageurl/mode.go b/internal/storageurl/mode.go index 759a29a93..f94c063ed 100644 --- a/internal/storageurl/mode.go +++ b/internal/storageurl/mode.go @@ -1,8 +1,13 @@ package storageurl import ( + "context" "fmt" + "os" "strings" + "sync" + + "github.com/Tencent/WeKnora/internal/logger" ) // Mode selects how stored files are referenced in an API response. @@ -20,8 +25,14 @@ const ( ModePublic Mode = "public" ) -// QueryParam is the request query parameter that selects the mode per call. -const QueryParam = "resource_urls" +const ( + // QueryParam is the request query parameter that selects the mode per call. + QueryParam = "resource_urls" + // EnvVar sets the deployment-wide default mode for requests that omit + // QueryParam. It sits alongside APP_EXTERNAL_URL, which is what actually + // makes `resource://` handles resolvable to a public `/r/` URL. + EnvVar = "RESOURCE_URL_MODE" +) // ParseMode validates a mode supplied by a client or by configuration. An empty // value yields ModeHandle so an unset parameter or setting keeps the default. @@ -38,3 +49,36 @@ func ParseMode(raw string) (Mode, error) { "invalid %s value %q: expected %q or %q", QueryParam, raw, ModeHandle, ModePublic) } } + +// badDefaultOnce keeps a misconfigured EnvVar to a single log line instead of one +// per request. The value itself is re-read every time so operators can roll a +// change out without a restart. +var badDefaultOnce sync.Once + +// DefaultMode returns the deployment-wide default from EnvVar. An unset or +// unparseable value yields ModeHandle, so a typo degrades to the safe default +// rather than failing every request. +func DefaultMode(ctx context.Context) Mode { + raw := strings.TrimSpace(os.Getenv(EnvVar)) + if raw == "" { + return ModeHandle + } + mode, err := ParseMode(raw) + if err != nil { + badDefaultOnce.Do(func() { + logger.Warnf(ctx, "ignoring %s: %v; falling back to %q", EnvVar, err, ModeHandle) + }) + return ModeHandle + } + return mode +} + +// ResolveMode combines the per-request query value with the deployment default. +// An explicit query value always wins; an invalid one is returned as an error so +// integrators find their typo instead of silently receiving handles. +func ResolveMode(ctx context.Context, queryValue string) (Mode, error) { + if strings.TrimSpace(queryValue) != "" { + return ParseMode(queryValue) + } + return DefaultMode(ctx), nil +} diff --git a/internal/storageurl/request.go b/internal/storageurl/request.go new file mode 100644 index 000000000..3ccc66da1 --- /dev/null +++ b/internal/storageurl/request.go @@ -0,0 +1,149 @@ +package storageurl + +import ( + "context" + + "github.com/Tencent/WeKnora/internal/types" + "github.com/Tencent/WeKnora/internal/types/interfaces" +) + +// NewRequestRewriter builds the Rewriter for one API request or response stream. +// +// ModeHandle yields a disabled Rewriter, so the default path resolves nothing +// and costs no access-grant rows. The tenant is taken from ctx because a +// reference may live on a tenant-configured storage backend rather than the +// process-wide default. +// +// No extra authorization gate applies here, unlike the `/files` proxy which +// rejects KB-restricted API keys. That gate exists because `/files` takes an +// arbitrary caller-supplied path it cannot bind to a KB allow-list. Here the +// references come from a response the caller is already authorized to receive, +// so the server — not the client — chooses which resources get a URL. +func NewRequestRewriter( + ctx context.Context, + mode Mode, + defaultSvc interfaces.FileService, + storageResolver interfaces.StorageBackendResolver, +) *Rewriter { + if mode != ModePublic { + return NewRewriter(nil, "API") + } + tenant, _ := types.TenantInfoFromContext(ctx) + var resolvers []interfaces.StorageBackendResolver + if storageResolver != nil { + resolvers = append(resolvers, storageResolver) + } + resolver := NewFileServiceResolver(tenant, defaultSvc, resolvers...).WithContext(ctx) + return NewRewriter(resolver, "API") +} + +// RewriteMessages replaces storage references in a message history response so +// clients receive loadable image URLs. Messages are mutated in place. +func (w *Rewriter) RewriteMessages(ctx context.Context, messages []*types.Message) { + if !w.Enabled() { + return + } + for _, message := range messages { + if message == nil { + continue + } + message.Content = w.String(ctx, message.Content) + for i := range message.Images { + message.Images[i].URL = w.Ref(ctx, message.Images[i].URL) + message.Images[i].Caption = w.String(ctx, message.Images[i].Caption) + } + message.KnowledgeReferences = w.CopyReferences(ctx, message.KnowledgeReferences) + w.rewriteAgentSteps(ctx, message.AgentSteps) + } +} + +// CopyReferences returns rewritten copies of retrieval results, covering both the +// chunk text and the structured image_info payload. +// +// Copies rather than in-place edits because an SSE references payload shares its +// *SearchResult pointers with the stream replay buffer and the assistant message +// being persisted; rewriting those in place would corrupt both. +func (w *Rewriter) CopyReferences(ctx context.Context, refs []*types.SearchResult) []*types.SearchResult { + if !w.Enabled() || refs == nil { + return refs + } + out := make([]*types.SearchResult, len(refs)) + for i, ref := range refs { + if ref == nil { + continue + } + rewritten := *ref + rewritten.Content = w.String(ctx, ref.Content) + rewritten.MatchedContent = w.String(ctx, ref.MatchedContent) + rewritten.ImageInfo = w.String(ctx, ref.ImageInfo) + out[i] = &rewritten + } + return out +} + +// CopyData returns a rewritten copy of an SSE metadata map, or data itself when +// it holds no storage reference. Agent tool results put renderable Markdown into +// this map, and its shape is tool-defined, so every string leaf is rewritten. +func (w *Rewriter) CopyData(ctx context.Context, data map[string]interface{}) map[string]interface{} { + if !w.Enabled() || data == nil { + return data + } + rewritten, changed := w.copyValue(ctx, data, 0) + if !changed { + return data + } + out, _ := rewritten.(map[string]interface{}) + return out +} + +// maxDataDepth bounds recursion into tool-defined metadata. Renderable content +// sits within a couple of levels; the cap only guards against a pathologically +// nested payload. +const maxDataDepth = 8 + +func (w *Rewriter) copyValue(ctx context.Context, value interface{}, depth int) (interface{}, bool) { + if depth > maxDataDepth { + return value, false + } + switch typed := value.(type) { + case string: + out := w.String(ctx, typed) + return out, out != typed + case map[string]interface{}: + out := make(map[string]interface{}, len(typed)) + changed := false + for key, item := range typed { + converted, itemChanged := w.copyValue(ctx, item, depth+1) + out[key] = converted + changed = changed || itemChanged + } + return out, changed + case []interface{}: + out := make([]interface{}, len(typed)) + changed := false + for i, item := range typed { + converted, itemChanged := w.copyValue(ctx, item, depth+1) + out[i] = converted + changed = changed || itemChanged + } + return out, changed + default: + return value, false + } +} + +// rewriteAgentSteps covers the reasoning trace, whose tool output embeds Markdown +// images for retrieved figures and generated charts. +func (w *Rewriter) rewriteAgentSteps(ctx context.Context, steps types.AgentSteps) { + for i := range steps { + steps[i].Thought = w.String(ctx, steps[i].Thought) + steps[i].ReasoningContent = w.String(ctx, steps[i].ReasoningContent) + for j := range steps[i].ToolCalls { + call := &steps[i].ToolCalls[j] + call.Reflection = w.String(ctx, call.Reflection) + if call.Result != nil { + call.Result.Output = w.String(ctx, call.Result.Output) + } + } + } +} diff --git a/internal/storageurl/stream.go b/internal/storageurl/stream.go index 0eec14545..449145948 100644 --- a/internal/storageurl/stream.go +++ b/internal/storageurl/stream.go @@ -100,6 +100,15 @@ func (s *StreamRewriter) Enabled() bool { return s != nil && s.rewriter.Enabled() } +// Rewriter exposes the underlying Rewriter for stream fields that arrive whole +// (references, metadata) and therefore need no holdback. +func (s *StreamRewriter) Rewriter() *Rewriter { + if s == nil { + return nil + } + return s.rewriter +} + // Push feeds the next chunk of the stream identified by key and returns the // rewritten content that is ready to emit. Set flush on the stream's terminal // chunk to release any held tail.