Files
ragflow/internal/agent/canvas/run_tracker.go
Hz_ 19b60132da fix(go-agent): use session IDs for cancellation and context flow (#17462)
## Summary

- Propagate request contexts through Agent Canvas execution and external
calls.
- Replace internal task IDs with session IDs while retaining `task_id`
as a wire alias.
- Complete session-scoped cancellation with Redis lease and token
validation.

## Testing

- Go backend tests passed.

<img width="1176" height="574" alt="image"
src="https://github.com/user-attachments/assets/b86560be-9b8d-45bb-97e9-921dffab8ebe"
/>
2026-07-28 14:59:34 +08:00

378 lines
14 KiB
Go

//
// Copyright 2026 The InfiniFlow Authors. All Rights Reserved.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
// run_tracker.go persists canvas-run business metadata to a Redis Hash.
// See plan §2.6 (Key 2: "agent:run:{run_id}"). This is the *business*
// channel — checkpoint payload (eino bytes) lives in checkpoint_store.go.
//
// Status code mapping (stored as int under the "status" field):
//
// 0 = running, 1 = succeeded, 2 = failed, 3 = cancelled, 4 = waiting_for_user.
package canvas
import (
"context"
"errors"
"strconv"
"time"
"github.com/redis/go-redis/v9"
redis2 "ragflow/internal/engine/redis"
)
// runKeyPrefix is the Redis Hash key namespace for run metadata.
// The full key is "agent:run:{run_id}".
const runKeyPrefix = "agent:run:"
// runStatus values for the "status" hash field.
const (
runStatusRunning = "0"
runStatusSucceeded = "1"
runStatusFailed = "2"
runStatusCancelled = "3"
runStatusWaiting = "4"
)
// runFieldInterruptID is the hash field that holds the eino interrupt id a
// run is paused on (plan §4.4, M3). Reuses the run hash — no dedicated key.
const runFieldInterruptID = "interrupt_id"
func runKey(runID string) string { return runKeyPrefix + runID }
const activeSessionKeyPrefix = "agent:active-session:"
const activeSessionLeaseTTL = 30 * time.Second
func activeSessionKey(sessionID string) string { return activeSessionKeyPrefix + sessionID }
// ActiveSession is the distributed ordinary-Agent run registration. Token is
// an internal lease owner token; it is never exposed as a business run id.
type ActiveSession struct {
SessionID string
Token string
UserID string
CanvasID string
RunID string
}
// RunTracker manages canvas-run metadata (canvas_id, status, checkpoint
// link, resume chain, ...) on a Redis Hash. Operations are explicit — the
// eino CheckPointStore does NOT write these fields, so callers (HTTP
// handler, cancel watcher) must invoke Start/Mark* at the right points.
type RunTracker struct {
client *redis.Client
ttl time.Duration
}
// NewRunTracker returns a tracker wired to the global Redis client. When
// the cache is uninitialized, client is nil; methods error in that case
// rather than panicking, and tests can inject a client via struct-literal
// construction.
func NewRunTracker(ttl time.Duration) *RunTracker {
var client *redis.Client
if rc := redis2.Get(); rc != nil {
client = rc.GetClient()
}
return &RunTracker{client: client, ttl: ttl}
}
// NewRunTrackerWithClient returns a tracker wired to a caller-supplied
// redis.Client. The intended use is tests that want to drive the
// RunTracker against an in-memory miniredis without touching the
// global Redis cache, but the helper is exported so non-test callers
// (multi-tenant deployments, custom Redis pools) can inject a
// dedicated client without going through the global cache singleton.
func NewRunTrackerWithClient(client *redis.Client, ttl time.Duration) *RunTracker {
return &RunTracker{client: client, ttl: ttl}
}
// Start records a new run as in-progress. canvasID and tenantID identify
// the source DSL and tenant; parentRunID may be empty for fresh runs and
// carries the source run-id for resume chains (R1 in plan §2.6).
//
// The HSet + Expire are sent through a pipeline so a TTL is set on the
// first write — without that, the key would have no expiry and a crashed
// run would leak the hash.
func (t *RunTracker) Start(ctx context.Context, runID, canvasID, tenantID, parentRunID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
now := time.Now().UnixMilli()
key := runKey(runID)
pipe := t.client.Pipeline()
pipe.HSet(ctx, key, map[string]any{
"canvas_id": canvasID,
"tenant_id": tenantID,
"parent_run_id": parentRunID,
"status": runStatusRunning,
"cancel_requested": 0,
"started_at": now,
})
pipe.Expire(ctx, key, t.ttl)
_, err := pipe.Exec(ctx)
return err
}
// AttachCheckpoint writes the latest checkpoint id for this run. It is the
// ONLY writer of the "checkpoint_id" field; every W1/W2/W3/W4 path (plan
// §2.6) must call this once before the run goroutine returns.
func (t *RunTracker) AttachCheckpoint(ctx context.Context, runID, checkpointID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
return t.client.HSet(ctx, runKey(runID), "checkpoint_id", checkpointID).Err()
}
// AttachInterrupt persists the eino interrupt id that paused this run
// (plan §4.4, M3). The id is written to the SAME hash key as the other run
// metadata — no new Redis key — so a process that crashes between "Invoke
// returned an interrupt" and "the next loop iteration resumes" can recover:
// the next process reads GetInterruptID at the top of its resume loop and
// calls compose.ResumeWithData with it. The ONLY writer of "interrupt_id".
func (t *RunTracker) AttachInterrupt(ctx context.Context, runID, interruptID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
return t.client.HSet(ctx, runKey(runID), runFieldInterruptID, interruptID).Err()
}
// GetInterruptID returns the persisted interrupt id for runID. The bool is
// false when no interrupt is pending (field absent or empty). A nil error
// with (id, false) means "no pending interrupt" — callers treat that as
// "fresh run, do a normal Invoke".
func (t *RunTracker) GetInterruptID(ctx context.Context, runID string) (string, bool, error) {
if t == nil || t.client == nil {
return "", false, errors.New("run tracker: redis client not initialized")
}
v, err := t.client.HGet(ctx, runKey(runID), runFieldInterruptID).Result()
if errors.Is(err, redis.Nil) {
return "", false, nil
}
if err != nil {
return "", false, err
}
if v == "" {
return "", false, nil
}
return v, true, nil
}
// ClearInterruptID removes the persisted interrupt id after a resumed run
// completes. Cancelling an unrelated turn does not discard the session's
// UserFillUp checkpoint.
func (t *RunTracker) ClearInterruptID(ctx context.Context, runID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
return t.client.HDel(ctx, runKey(runID), runFieldInterruptID).Err()
}
// MarkSucceeded transitions the run to status=1 and stamps finished_at.
func (t *RunTracker) MarkSucceeded(ctx context.Context, runID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
return t.client.HSet(ctx, runKey(runID),
"status", runStatusSucceeded,
"finished_at", time.Now().UnixMilli(),
).Err()
}
// MarkFailed transitions the run to status=2 and records the reason.
func (t *RunTracker) MarkFailed(ctx context.Context, runID, reason string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
return t.client.HSet(ctx, runKey(runID),
"status", runStatusFailed,
"finished_at", time.Now().UnixMilli(),
"failure_reason", reason,
).Err()
}
// MarkCancelled transitions the run to status=3 and sets the cancel flag.
func (t *RunTracker) MarkCancelled(ctx context.Context, runID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
key := runKey(runID)
pipe := t.client.Pipeline()
pipe.HSet(ctx, key,
"status", runStatusCancelled,
"finished_at", time.Now().UnixMilli(),
"cancel_requested", 1,
)
pipe.Expire(ctx, key, t.ttl)
_, err := pipe.Exec(ctx)
return err
}
// RegisterActiveSession atomically acquires the distributed active-run lease
// for sessionID. A cancel marker found without an active lease belongs to an
// older run and is removed before the new owner is registered. A false result
// means another instance already owns the session.
func (t *RunTracker) RegisterActiveSession(ctx context.Context, active ActiveSession) (bool, error) {
if t == nil || t.client == nil {
return false, errors.New("run tracker: redis client not initialized")
}
const script = `
if redis.call("EXISTS", KEYS[1]) == 1 then return 0 end
redis.call("DEL", KEYS[2])
redis.call("HSET", KEYS[1],
"session_id", ARGV[1], "token", ARGV[2], "user_id", ARGV[3],
"canvas_id", ARGV[4], "run_id", ARGV[5], "started_at", ARGV[6])
redis.call("PEXPIRE", KEYS[1], ARGV[7])
return 1`
result, err := t.client.Eval(ctx, script,
[]string{activeSessionKey(active.SessionID), cancelKey(active.SessionID)},
active.SessionID, active.Token, active.UserID, active.CanvasID, active.RunID,
strconv.FormatInt(time.Now().UnixMilli(), 10), strconv.FormatInt(t.leaseTTL().Milliseconds(), 10),
).Int64()
return result == 1, err
}
// GetActiveSession returns the distributed active run for sessionID. A nil
// result means the session is not running.
func (t *RunTracker) GetActiveSession(ctx context.Context, sessionID string) (*ActiveSession, error) {
if t == nil || t.client == nil {
return nil, errors.New("run tracker: redis client not initialized")
}
fields, err := t.client.HGetAll(ctx, activeSessionKey(sessionID)).Result()
if err != nil || len(fields) == 0 {
return nil, err
}
return &ActiveSession{
SessionID: fields["session_id"],
Token: fields["token"],
UserID: fields["user_id"],
CanvasID: fields["canvas_id"],
RunID: fields["run_id"],
}, nil
}
// RefreshActiveSession extends the lease only while token still owns it.
func (t *RunTracker) RefreshActiveSession(ctx context.Context, sessionID, token string) (bool, error) {
if t == nil || t.client == nil {
return false, errors.New("run tracker: redis client not initialized")
}
const script = `if redis.call("HGET", KEYS[1], "token") == ARGV[1] then return redis.call("PEXPIRE", KEYS[1], ARGV[2]) end return 0`
result, err := t.client.Eval(ctx, script, []string{activeSessionKey(sessionID)}, token,
strconv.FormatInt(t.leaseTTL().Milliseconds(), 10)).Int64()
return result == 1, err
}
// ReleaseActiveSession deletes the registration and cancel marker only when
// token still owns the session, preventing stale cleanup from touching a newer
// run.
func (t *RunTracker) ReleaseActiveSession(ctx context.Context, sessionID, token string) (bool, error) {
if t == nil || t.client == nil {
return false, errors.New("run tracker: redis client not initialized")
}
const script = `
if redis.call("HGET", KEYS[1], "token") ~= ARGV[1] then return 0 end
redis.call("DEL", KEYS[2])
redis.call("DEL", KEYS[1])
return 1`
result, err := t.client.Eval(ctx, script,
[]string{activeSessionKey(sessionID), cancelKey(sessionID)}, token).Int64()
return result == 1, err
}
// RequestCancelActiveSession publishes a cancel marker only while token still
// owns the active session lease. This prevents a stale owner from cancelling
// a newer run after its lease has been replaced or released.
func (t *RunTracker) RequestCancelActiveSession(ctx context.Context, sessionID, token string) (bool, error) {
if t == nil || t.client == nil {
return false, errors.New("run tracker: redis client not initialized")
}
const script = `
if redis.call("HGET", KEYS[1], "token") ~= ARGV[1] then return 0 end
redis.call("SET", KEYS[2], "x", "PX", ARGV[2])
return 1`
result, err := t.client.Eval(ctx, script,
[]string{activeSessionKey(sessionID), cancelKey(sessionID)}, token,
strconv.FormatInt(RequestCancelTTL.Milliseconds(), 10)).Int64()
return result == 1, err
}
// WatchActiveSession keeps the distributed lease alive for the lifetime of
// ctx. It invokes onLost if another owner replaces the lease or Redis cannot
// confirm ownership before the current lease can expire.
func (t *RunTracker) WatchActiveSession(ctx context.Context, sessionID, token string, onLost func()) {
if t == nil || t.client == nil {
return
}
leaseTTL := t.leaseTTL()
interval := leaseTTL / 3
if interval > 10*time.Second {
interval = 10 * time.Second
}
if interval < 10*time.Millisecond {
interval = 10 * time.Millisecond
}
ticker := time.NewTicker(interval)
defer ticker.Stop()
lastConfirmed := time.Now()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
owned, err := t.RefreshActiveSession(ctx, sessionID, token)
switch {
case err == nil && owned:
lastConfirmed = time.Now()
case err == nil && !owned, time.Since(lastConfirmed) >= leaseTTL*2/3:
if onLost != nil {
onLost()
}
return
}
}
}
}
func (t *RunTracker) leaseTTL() time.Duration {
if t.ttl > 0 && t.ttl < activeSessionLeaseTTL {
return t.ttl
}
return activeSessionLeaseTTL
}
// MarkWaiting records a normal UserFillUp pause. It is not a failure.
func (t *RunTracker) MarkWaiting(ctx context.Context, runID string) error {
if t == nil || t.client == nil {
return errors.New("run tracker: redis client not initialized")
}
key := runKey(runID)
pipe := t.client.Pipeline()
pipe.HSet(ctx, key, "status", runStatusWaiting)
pipe.Expire(ctx, key, t.ttl)
_, err := pipe.Exec(ctx)
return err
}
// Get returns all hash fields for a run. The empty map (not nil) plus a
// nil error means "no such run" — callers can detect this with len(map)==0
// if they need to distinguish from a key that exists with no fields.
func (t *RunTracker) Get(ctx context.Context, runID string) (map[string]string, error) {
if t == nil || t.client == nil {
return nil, errors.New("run tracker: redis client not initialized")
}
return t.client.HGetAll(ctx, runKey(runID)).Result()
}