mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-08-31 00:50:02 +08:00
feat(api): add resource_urls=public to return directly loadable file URLs
API responses reference stored files as opaque resource:// handles, so an integrating app had to make a second authenticated call to /files for every image before it could render anything. Add an opt-in that resolves those references server-side into time-limited HTTP(S) URLs, using the same mechanism the IM channels already rely on: - per request: ?resource_urls=public (default: handle, unchanged) - per deployment: RESOURCE_URL_MODE=public Applied to the chat SSE endpoints (knowledge-chat, agent-chat, continue-stream), message history load, and knowledge-search. Streamed answers buffer a trailing incomplete reference so a handle split across two deltas is still rewritten. References that cannot become an HTTP URL (for example local storage with no APP_EXTERNAL_URL) stay handles, so clients keep the /files fallback. Embed channels are deliberately excluded: their visitors are anonymous. SSE payloads share SearchResult pointers and metadata maps with the stream replay buffer and the message being persisted, so those are rewritten as copies. Co-authored-by: lyingbug <lyingbug@users.noreply.github.com>
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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/<token>` 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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user