Merge pull request #80 from dj-oyu/feat/upstream-test-compat

feat: implement upstream event/steering/interrupt features
This commit is contained in:
dj-oyu 2026-03-26 01:28:17 +09:00 committed by GitHub
commit 9d9a986ccd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 234 additions and 16 deletions

View file

@ -479,12 +479,13 @@ func (al *AgentLoop) llmWorker(ctx context.Context, queue <-chan bus.InboundMess
continue
}
al.llmWorkerNormal(ctx, msg)
al.llmWorkerNormal(ctx, msg, queue)
}
}
// llmWorkerNormal processes a single non-PDF message.
func (al *AgentLoop) llmWorkerNormal(ctx context.Context, msg bus.InboundMessage) {
// queue is optional (nil when called outside Run loop).
func (al *AgentLoop) llmWorkerNormal(ctx context.Context, msg bus.InboundMessage, queue <-chan bus.InboundMessage) {
al.activeRequests.Add(1)
defer al.activeRequests.Done()
@ -508,9 +509,66 @@ func (al *AgentLoop) llmWorkerNormal(ctx context.Context, msg bus.InboundMessage
response = fmt.Sprintf("Error processing message: %v", err)
}
// Auto-continue: if inbound messages arrived in the queue while this
// turn was running, treat them as steering continuations so the agent
// responds with the full context instead of two separate turns.
if queue != nil {
if drained := al.drainQueueAsSteering(queue, msg); len(drained) > 0 {
agent := al.agentForSession(msg.SessionKey)
if agent == nil {
agent = al.registry.GetDefaultAgent()
}
if agent != nil {
sessionKey := msg.SessionKey
if sessionKey == "" {
route, _, _ := al.resolveMessageRoute(msg)
sessionKey = resolveScopeKey(route, msg.SessionKey)
}
contResp, contErr := al.continueWithSteeringMessages(
ctx, agent, sessionKey, msg.Channel, msg.ChatID, drained,
)
if contErr == nil && contResp != "" {
response = contResp
}
}
}
}
al.sendResponseIfNeeded(ctx, msg, response)
}
// drainQueueAsSteering non-blocking drains pending messages from the
// llmQueue that belong to the same chat as the original message and
// returns them as steering-style provider messages.
func (al *AgentLoop) drainQueueAsSteering(
queue <-chan bus.InboundMessage,
orig bus.InboundMessage,
) []providers.Message {
var msgs []providers.Message
for {
select {
case m, ok := <-queue:
if !ok {
return msgs
}
if m.Channel == orig.Channel && m.ChatID == orig.ChatID {
msgs = append(msgs, providers.Message{
Role: "user",
Content: m.Content,
})
} else {
// Re-queue by pushing to steering for later processing
al.enqueueSteeringMessage("", "", providers.Message{
Role: "user",
Content: m.Content,
})
}
default:
return msgs
}
}
}
// llmWorkerPDF handles a bare-PDF message with two-phase follow-up collection.
func (al *AgentLoop) llmWorkerPDF(ctx context.Context, msg bus.InboundMessage, queue <-chan bus.InboundMessage) {
// Phase 1: wait for OCR keywords (figures/図版) — up to 5 seconds.
@ -553,13 +611,13 @@ func (al *AgentLoop) llmWorkerPDF(ctx context.Context, msg bus.InboundMessage, q
Content: followUpText,
Metadata: msg.Metadata,
}
al.llmWorkerNormal(ctx, followUpMsg)
al.llmWorkerNormal(ctx, followUpMsg, nil)
}
// Re-queue messages from other chats that were buffered.
otherMsgs := extractNonChatMessages(buffered, msg.ChatID)
for _, other := range otherMsgs {
al.llmWorkerNormal(ctx, other)
al.llmWorkerNormal(ctx, other, nil)
}
}
@ -1179,7 +1237,31 @@ func (al *AgentLoop) ProcessDirectWithChannel(
},
}
return al.processMessage(ctx, msg)
response, err := al.processMessage(ctx, msg)
if err != nil {
return response, err
}
// If steering messages arrived during the LLM call, the initial
// response is stale. Discard it and do a continuation turn that
// includes the steering messages for a fresh response.
steeringMsgs := al.dequeueSteeringMessagesForScopeWithFallback(sessionKey)
if len(steeringMsgs) > 0 {
agent := al.agentForSession(sessionKey)
if agent == nil {
agent = al.registry.GetDefaultAgent()
}
if agent != nil {
contResp, contErr := al.continueWithSteeringMessages(
ctx, agent, sessionKey, channel, chatID, steeringMsgs,
)
if contErr == nil && contResp != "" {
return contResp, nil
}
}
}
return response, nil
}
// ProcessHeartbeat processes a heartbeat request without session history.

View file

@ -62,6 +62,14 @@ func (al *AgentLoop) buildCommandsRuntime(agent *AgentInstance, sessionKey strin
}
old := agent.Model
agent.Model = value
// Rebuild candidates so subsequent LLM calls use the new
// model's provider/endpoint instead of the old one.
agent.Candidates = resolveModelCandidates(
cfg,
cfg.Agents.Defaults.Provider,
value,
nil,
)
return old, nil
},
SwitchChannel: func(value string) error {

View file

@ -33,6 +33,10 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt
scope := al.newTurnEventScope(agent.ID, opts.SessionKey)
turnStart := time.Now()
// Create a cancelable context for hard abort support
turnCtx, turnCancelFn := context.WithCancel(ctx)
defer turnCancelFn()
// Register a turnState so the interrupt API can find this turn
ts := &turnState{
turnID: scope.turnID,
@ -44,6 +48,12 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt
phase: TurnPhaseRunning,
startedAt: turnStart,
agent: agent,
turnCancel: turnCancelFn,
}
// Bind session store and capture initial history length for rollback
if agent.Sessions != nil {
ts.session = agent.Sessions
ts.initialHistoryLength = len(agent.Sessions.GetHistory(opts.SessionKey))
}
al.registerActiveTurn(ts)
defer al.clearActiveTurn(ts)
@ -65,7 +75,7 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt
// -0. Create cancelable child context and register active task
taskCtx, taskCancel := context.WithCancel(ctx)
taskCtx, taskCancel := context.WithCancel(turnCtx)
defer taskCancel()
@ -354,7 +364,9 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt
historyMsg = opts.HistoryMessage
}
if historyMsg != "" {
agent.Sessions.AddMessage(opts.SessionKey, "user", historyMsg)
}
// 4. Record user prompt for stats
@ -413,6 +425,30 @@ func (al *AgentLoop) runAgentLoop(ctx context.Context, agent *AgentInstance, opt
finalContent, iteration, err = al.runLLMIteration(ctx, agent, messages, opts, task, curPlanStatus, scope)
if err != nil {
// Check for hard abort: if the turn was hard-aborted, restore
// session history and return empty response.
ts.mu.RLock()
isHardAbort := ts.hardAbort
ts.mu.RUnlock()
if isHardAbort {
if ts.session != nil {
history := ts.session.GetHistory(opts.SessionKey)
if ts.initialHistoryLength < len(history) {
ts.session.SetHistory(opts.SessionKey, history[:ts.initialHistoryLength])
}
}
abortMeta := EventMeta{
AgentID: agent.ID,
TurnID: scope.turnID,
SessionKey: opts.SessionKey,
Iteration: iteration,
}
al.emitEvent(EventKindTurnEnd, abortMeta, TurnEndPayload{
Status: TurnEndStatusAborted,
Duration: time.Since(turnStart),
})
return "", nil
}
return "", err
}
@ -868,10 +904,24 @@ func (al *AgentLoop) runLLMIteration(
// Also inject any initial steering messages from the process options
if len(opts.InitialSteeringMessages) > 0 {
messages = append(messages, opts.InitialSteeringMessages...)
// Persist original refs in session history
for _, sm := range opts.InitialSteeringMessages {
agent.Sessions.AddFullMessage(opts.SessionKey, sm)
}
// Resolve media refs for the provider call while keeping
// the originals in session history with raw refs.
cfg := al.GetConfig()
maxMedia := cfg.Agents.Defaults.GetMaxMediaSize()
resolved := resolveMediaRefs(opts.InitialSteeringMessages, al.mediaStore, maxMedia)
messages = append(messages, resolved...)
}
for iteration < agent.MaxIterations {
// Check for context cancellation (e.g. hard abort) before each iteration
if ctx.Err() != nil {
return "", iteration, ctx.Err()
}
iteration++
if msg := hooks.OnIterationStart(iteration); msg != "" {
@ -1014,7 +1064,16 @@ func (al *AgentLoop) runLLMIteration(
HasReasoning: response.Reasoning != "",
})
go al.handleReasoning(ctx, response.Reasoning, opts.Channel, al.targetReasoningChannelID(opts.Channel))
reasoningText := response.Reasoning
if reasoningText == "" {
reasoningText = response.ReasoningContent
}
go al.handleReasoning(
context.WithoutCancel(ctx),
reasoningText,
opts.Channel,
al.targetReasoningChannelID(opts.Channel),
)
logger.DebugCF("agent", "LLM response",
map[string]any{
@ -1090,12 +1149,36 @@ func (al *AgentLoop) runLLMIteration(
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
// Execute tool calls and collect results
lastBlocker := al.executeToolCalls(ctx, agent, normalizedToolCalls, &messages, opts, hooks, iteration, scope)
execResult := al.executeToolCalls(ctx, agent, normalizedToolCalls, &messages, opts, hooks, iteration, scope)
// Graceful interrupt: make a terminal LLM call with no tool
// definitions so the agent can produce a final summary.
if execResult.gracefulInterrupt {
hintMsg := "Interrupt requested. Stop scheduling tools and provide a short final summary."
if execResult.gracefulHint != "" {
hintMsg += "\n\nInterrupt hint: " + execResult.gracefulHint
}
messages = append(messages, providers.Message{
Role: "user",
Content: hintMsg,
})
terminalResp, termErr := al.callLLMWithRetry(
ctx, agent, &messages, opts,
nil, // no tool definitions
al.selectCandidates(agent, "", messages),
agent.Model, nil, iteration+1, scope,
)
if termErr == nil && terminalResp != nil {
finalContent = terminalResp.Content
}
break
}
// Tick TTL-based tool expiry after execution
agent.Tools.TickTTL()
hooks.InjectReminders(iteration, &messages, lastBlocker)
hooks.InjectReminders(iteration, &messages, execResult.lastBlocker)
hooks.RefreshSystemPrompt(messages)
}
@ -1107,6 +1190,13 @@ func (al *AgentLoop) runLLMIteration(
return finalContent, iteration, nil
}
// toolExecResult holds results from executeToolCalls.
type toolExecResult struct {
lastBlocker string
gracefulInterrupt bool
gracefulHint string
}
// executeToolCalls runs each tool call sequentially, publishes results,
// and returns the last blocker (error content) for reminder injection.
func (al *AgentLoop) executeToolCalls(
@ -1118,12 +1208,26 @@ func (al *AgentLoop) executeToolCalls(
hooks iterationHooks,
iteration int,
scope turnEventScope,
) string {
var lastBlocker string
) toolExecResult {
var result toolExecResult
steered := false
gracefulSkip := false
for i, tc := range toolCalls {
// Check for graceful interrupt between tool calls
if !gracefulSkip {
if ts := al.getActiveTurnState(opts.SessionKey); ts != nil {
ts.mu.RLock()
if ts.gracefulInterrupt {
gracefulSkip = true
result.gracefulInterrupt = true
result.gracefulHint = ts.gracefulInterruptHint
}
ts.mu.RUnlock()
}
}
// Check for pending steering messages between tool calls
if i > 0 && !steered {
if i > 0 && !steered && !gracefulSkip {
steeringMsgs := al.dequeueSteeringMessagesForScopeWithFallback(opts.SessionKey)
if len(steeringMsgs) > 0 {
steered = true
@ -1149,6 +1253,16 @@ func (al *AgentLoop) executeToolCalls(
}
}
// Skip remaining tools if graceful interrupt was requested
if gracefulSkip {
*messages = append(*messages, providers.Message{
Role: "tool",
Content: "Skipped due to graceful interrupt.",
ToolCallID: tc.ID,
})
continue
}
// Skip remaining tools if steering was injected
if steered {
skipMeta := EventMeta{
@ -1312,7 +1426,7 @@ func (al *AgentLoop) executeToolCalls(
contentForLLM = toolResult.Err.Error()
}
if toolResult.IsError || toolResult.Err != nil {
lastBlocker = contentForLLM
result.lastBlocker = contentForLLM
}
toolResultMsg := providers.Message{
@ -1323,7 +1437,21 @@ func (al *AgentLoop) executeToolCalls(
*messages = append(*messages, toolResultMsg)
agent.Sessions.AddFullMessage(opts.SessionKey, toolResultMsg)
}
return lastBlocker
// Re-check graceful interrupt after all tool execution (may have been
// set during the last tool's execution).
if !result.gracefulInterrupt {
if ts := al.getActiveTurnState(opts.SessionKey); ts != nil {
ts.mu.RLock()
if ts.gracefulInterrupt {
result.gracefulInterrupt = true
result.gracefulHint = ts.gracefulInterruptHint
}
ts.mu.RUnlock()
}
}
return result
}
// forceTextResponse makes a final LLM call without tools when max iterations