Files
AIRouter/internal/context/assembler.go
T
freedakgmail 93a469061d
CI / lint (push) Has been cancelled
CI / test (push) Has been cancelled
CI / build (push) Has been cancelled
CI / security-scan (push) Has been cancelled
初始提交:边缘AI算力机统一AI通讯层
2026-08-03 07:44:05 +08:00

209 lines
5.5 KiB
Go

package context
import (
"github.com/edgeai/gateway/internal/config"
"github.com/edgeai/gateway/pkg/api"
)
// Assembler assembles context messages for a chat request.
type Assembler struct {
estimator *TokenEstimator
cfg *config.ContextConfig
}
// NewAssembler creates a new context assembler.
func NewAssembler(cfg *config.ContextConfig) *Assembler {
return &Assembler{
estimator: NewTokenEstimator(),
cfg: cfg,
}
}
// AssembleResult contains the assembled messages and metadata.
type AssembleResult struct {
Messages []api.Message
InputTokens int
Trimmed bool
TrimmedCount int
}
// Assemble combines session history with new messages, applying context window limits.
func (a *Assembler) Assemble(history []api.Message, newMessages []api.Message, contextWindow int, maxOutputTokens int, policy string) *AssembleResult {
// Calculate available context for history
availableForHistory := contextWindow - maxOutputTokens
if availableForHistory < 0 {
availableForHistory = contextWindow / 2
}
// Apply safety margin
availableForHistory = int(float64(availableForHistory) * (1.0 - a.cfg.SafetyMarginRatio))
// Combine all messages
allMessages := make([]api.Message, 0, len(history)+len(newMessages))
allMessages = append(allMessages, history...)
allMessages = append(allMessages, newMessages...)
// Estimate total tokens
totalTokens := a.estimateAllTokens(allMessages)
if totalTokens <= availableForHistory {
return &AssembleResult{
Messages: allMessages,
InputTokens: totalTokens,
Trimmed: false,
}
}
// Need to trim — apply policy
trimmed := a.applyPolicy(allMessages, availableForHistory, policy)
return &AssembleResult{
Messages: trimmed.messages,
InputTokens: trimmed.tokens,
Trimmed: true,
TrimmedCount: len(allMessages) - len(trimmed.messages),
}
}
type trimResult struct {
messages []api.Message
tokens int
}
func (a *Assembler) applyPolicy(messages []api.Message, budget int, policy string) trimResult {
switch policy {
case "recent_only":
return a.trimRecentOnly(messages, budget)
case "summary_and_recent":
return a.trimSummaryAndRecent(messages, budget)
case "full":
return a.trimFull(messages, budget)
default:
return a.trimSummaryAndRecent(messages, budget)
}
}
// trimRecentOnly keeps only the most recent messages within budget.
func (a *Assembler) trimRecentOnly(messages []api.Message, budget int) trimResult {
result := make([]api.Message, 0)
tokens := 0
// Iterate from the end (most recent first)
for i := len(messages) - 1; i >= 0; i-- {
msgTokens := a.estimateMsgTokens(messages[i])
if tokens+msgTokens > budget && len(result) > 0 {
break
}
// Prepend to maintain order
result = append([]api.Message{messages[i]}, result...)
tokens += msgTokens
}
return trimResult{messages: result, tokens: tokens}
}
// trimSummaryAndRecent keeps system message + a summary placeholder + recent messages.
func (a *Assembler) trimSummaryAndRecent(messages []api.Message, budget int) trimResult {
if len(messages) == 0 {
return trimResult{}
}
// Always keep system messages at the front
systemMsgs := []api.Message{}
rest := []api.Message{}
for _, m := range messages {
if m.Role == "system" {
systemMsgs = append(systemMsgs, m)
} else {
rest = append(rest, m)
}
}
systemTokens := 0
for _, m := range systemMsgs {
systemTokens += a.estimateMsgTokens(m)
}
// Reserve space for a summary placeholder (~50 tokens)
summaryTokens := 50
availableForRecent := budget - systemTokens - summaryTokens
if availableForRecent < 0 {
availableForRecent = budget / 2
}
// Keep most recent messages
recentMsgs := []api.Message{}
recentTokens := 0
for i := len(rest) - 1; i >= 0; i-- {
msgTokens := a.estimateMsgTokens(rest[i])
if recentTokens+msgTokens > availableForRecent && len(recentMsgs) > 0 {
break
}
recentMsgs = append([]api.Message{rest[i]}, recentMsgs...)
recentTokens += msgTokens
}
// Add summary placeholder if we trimmed anything
result := make([]api.Message, 0, len(systemMsgs)+1+len(recentMsgs))
result = append(result, systemMsgs...)
if len(recentMsgs) < len(rest) {
result = append(result, api.Message{
Role: "system",
Content: "[Earlier conversation history has been summarized and omitted.]",
})
}
result = append(result, recentMsgs...)
return trimResult{
messages: result,
tokens: systemTokens + summaryTokens + recentTokens,
}
}
// trimFull keeps messages as-is but truncates the oldest if over budget.
func (a *Assembler) trimFull(messages []api.Message, budget int) trimResult {
result := make([]api.Message, 0, len(messages))
tokens := 0
// Keep system messages, trim oldest non-system messages
systemMsgs := []api.Message{}
rest := []api.Message{}
for _, m := range messages {
if m.Role == "system" {
systemMsgs = append(systemMsgs, m)
} else {
rest = append(rest, m)
}
}
for _, m := range systemMsgs {
t := a.estimateMsgTokens(m)
tokens += t
result = append(result, m)
}
for _, m := range rest {
t := a.estimateMsgTokens(m)
if tokens+t > budget {
break
}
tokens += t
result = append(result, m)
}
return trimResult{messages: result, tokens: tokens}
}
func (a *Assembler) estimateAllTokens(messages []api.Message) int {
total := 0
for _, m := range messages {
total += a.estimateMsgTokens(m)
}
return total
}
func (a *Assembler) estimateMsgTokens(msg api.Message) int {
content, _ := msg.Content.(string)
return a.estimator.EstimateText(msg.Role) + a.estimator.EstimateText(content) + 4
}