Go: add context to redis client (#17689)

Signed-off-by: Jin Hai <haijin.chn@gmail.com>
This commit is contained in:
Jin Hai
2026-08-02 22:50:54 +08:00
committed by GitHub
parent b2521ebf51
commit 266837eb33
31 changed files with 302 additions and 308 deletions

View File

@@ -129,13 +129,13 @@ func (h *AgentHandler) WithDocumentService(s documentAccessChecker) *AgentHandle
// NewAgentHandler create agent handler
func NewAgentHandler(agentService *service.AgentService, fileService *file.FileService) *AgentHandler {
func NewAgentHandler(ctx context.Context, agentService *service.AgentService, fileService *file.FileService) *AgentHandler {
return &AgentHandler{
agentService: agentService,
chatRunner: agentService,
fileService: fileService,
loader: agentService,
redisGet: func(key string) (string, error) { return redis.Get().Get(key) },
redisGet: func(key string) (string, error) { return redis.Get().Get(ctx, key) },
redisStore: redis.Get(),
newExecutor: func(taskCtx *task.TaskContext, canvasID string, docBulkSize int) (debugExecutor, error) {
return task.NewPipelineExecutor(taskCtx, canvasID, docBulkSize)

View File

@@ -113,6 +113,7 @@ func TestParseAgentLogs(t *testing.T) {
func TestGetAgentLogs_E2EViaMiniredis(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := t.Context()
db := setupHandlerAgentsTestDB(t)
orig := dao.DB
dao.DB = db
@@ -134,13 +135,13 @@ func TestGetAgentLogs_E2EViaMiniredis(t *testing.T) {
`{"component_id":"File","trace":[{"progress":1,"message":"parsed","datetime":"10:00:00","timestamp":1.0,"elapsed_time":0}]},` +
`{"component_id":"END","trace":[{"progress":1,"message":"done","datetime":"10:00:01","timestamp":2.0,"elapsed_time":1.0}]}` +
`]`
if err := rdb.Set(context.Background(), logKey, arrayPayload, 0).Err(); err != nil {
if err = rdb.Set(ctx, logKey, arrayPayload, 0).Err(); err != nil {
t.Fatalf("seed redis: %v", err)
}
h := NewAgentHandler(service.NewAgentService(), nil).
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithRedisGetter(func(key string) (string, error) {
return rdb.Get(context.Background(), key).Result()
return rdb.Get(ctx, key).Result()
})
run := func(messageID string) map[string]interface{} {
@@ -224,6 +225,7 @@ func clientConsidersComplete(arr []map[string]interface{}) bool {
// byte-for-byte through JSON round-tripping.
func TestGetAgentLogs_EndSignalCompletion(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := t.Context()
db := setupHandlerAgentsTestDB(t)
orig := dao.DB
@@ -246,7 +248,7 @@ func TestGetAgentLogs_EndSignalCompletion(t *testing.T) {
`{"component_id":"File","trace":[{"progress":1,"message":"parsed","datetime":"10:00:00","timestamp":1.0,"elapsed_time":0}]},` +
`{"component_id":"END","trace":[{"progress":1,"message":"run finished","datetime":"10:00:01","timestamp":2.0,"elapsed_time":1.0}]}` +
`]`
if err := rdb.Set(context.Background(), "c1-msg-good-logs", goodPayload, 0).Err(); err != nil {
if err = rdb.Set(ctx, "c1-msg-good-logs", goodPayload, 0).Err(); err != nil {
t.Fatalf("seed redis: %v", err)
}
@@ -257,13 +259,13 @@ func TestGetAgentLogs_EndSignalCompletion(t *testing.T) {
`{"component_id":"File","trace":[{"progress":1,"message":"parsed","datetime":"10:00:00","timestamp":1.0,"elapsed_time":0}]},` +
`{"component_id":"END","trace":[{"progress":1,"message":"","datetime":"10:00:01","timestamp":2.0,"elapsed_time":1.0}]}` +
`]`
if err := rdb.Set(context.Background(), "c1-msg-bad-logs", badPayload, 0).Err(); err != nil {
if err = rdb.Set(ctx, "c1-msg-bad-logs", badPayload, 0).Err(); err != nil {
t.Fatalf("seed redis: %v", err)
}
h := NewAgentHandler(service.NewAgentService(), nil).
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithRedisGetter(func(key string) (string, error) {
return rdb.Get(context.Background(), key).Result()
return rdb.Get(ctx, key).Result()
})
call := func(messageID string) []map[string]interface{} {
@@ -310,7 +312,7 @@ type capturedStore struct {
data map[string]string
}
func (s *capturedStore) Set(key, value string, _ time.Duration) bool {
func (s *capturedStore) Set(ctx context.Context, key, value string, _ time.Duration) bool {
s.mu.Lock()
defer s.mu.Unlock()
if s.data == nil {
@@ -320,7 +322,7 @@ func (s *capturedStore) Set(key, value string, _ time.Duration) bool {
return true
}
func (s *capturedStore) get(key string) (string, bool) {
func (s *capturedStore) Get(ctx context.Context, key string) (string, bool) {
s.mu.Lock()
defer s.mu.Unlock()
v, ok := s.data[key]
@@ -463,10 +465,10 @@ func TestRespondWithDebugResult_ErrorCarriesMessageID(t *testing.T) {
// before — the failure log was written but unreachable because message_id was
// dropped on the error path.
func TestRunCanvasPipelineDebug_ErrorStillExposesMessageID(t *testing.T) {
ctx := context.Background()
ctx := t.Context()
store := &capturedStore{}
h := NewAgentHandler(service.NewAgentService(), nil).
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithRedisStore(store).
WithNewExecutor(func(taskCtx *task.TaskContext, canvasID string, docBulkSize int) (debugExecutor, error) {
return &fakeDebugExecutor{
@@ -489,7 +491,7 @@ func TestRunCanvasPipelineDebug_ErrorStillExposesMessageID(t *testing.T) {
// The failure log must be written under the composed key.
key := "c1-" + result.MessageID + "-logs"
raw, ok := store.get(key)
raw, ok := store.Get(ctx, key)
if !ok {
t.Fatalf("failure log not written under key %q; store keys=%v", key, keysOf(store))
}
@@ -511,10 +513,10 @@ func TestRunCanvasPipelineDebug_ErrorStillExposesMessageID(t *testing.T) {
// satisfy the completion predicate (END last, non-empty END message) so the
// Log box stops polling.
func TestRunCanvasPipelineDebug_WiresMessageIDAndLog(t *testing.T) {
ctx := context.Background()
ctx := t.Context()
store := &capturedStore{}
h := NewAgentHandler(service.NewAgentService(), nil).
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithRedisStore(store).
WithNewExecutor(func(taskCtx *task.TaskContext, canvasID string, docBulkSize int) (debugExecutor, error) {
return &fakeDebugExecutor{
@@ -537,7 +539,7 @@ func TestRunCanvasPipelineDebug_WiresMessageIDAndLog(t *testing.T) {
// The log array must be written under the composed key.
key := "c1-" + result.MessageID + "-logs"
raw, ok := store.get(key)
raw, ok := store.Get(ctx, key)
if !ok {
t.Fatalf("log not written under key %q; store keys=%v", key, keysOf(store))
}
@@ -570,8 +572,8 @@ type miniredisDebugStore struct {
rdb *goredis.Client
}
func (s miniredisDebugStore) Set(key, value string, ttl time.Duration) bool {
if err := s.rdb.Set(context.Background(), key, value, ttl).Err(); err != nil {
func (s miniredisDebugStore) Set(ctx context.Context, key, value string, ttl time.Duration) bool {
if err := s.rdb.Set(ctx, key, value, ttl).Err(); err != nil {
return false
}
return true
@@ -589,7 +591,7 @@ func (s miniredisDebugStore) Set(key, value string, ttl time.Duration) bool {
// seeds Redis directly rather than going through the writer.
func TestRunCanvasPipelineDebug_WriteThenReadViaMiniredis(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := context.Background()
ctx := t.Context()
db := setupHandlerAgentsTestDB(t)
orig := dao.DB
@@ -608,7 +610,7 @@ func TestRunCanvasPipelineDebug_WriteThenReadViaMiniredis(t *testing.T) {
// Both seams point at the same miniredis: the writer stores via
// WithRedisStore and the reader fetches via WithRedisGetter, mirroring
// production where both hit one Redis.
h := NewAgentHandler(service.NewAgentService(), nil).
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithRedisStore(miniredisDebugStore{rdb: rdb}).
WithRedisGetter(func(key string) (string, error) {
return rdb.Get(ctx, key).Result()

View File

@@ -112,7 +112,8 @@ func TestListAgentVersionsHandler_Success(t *testing.T) {
},
})
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.ListVersions(c)
if w.Code != http.StatusOK {
@@ -162,7 +163,8 @@ func TestListAgentVersionsHandler_NoPermission(t *testing.T) {
// Canvas owned by user-b
db.Create(&entity.UserCanvas{ID: "canvas-b", UserID: "user-b", Title: sptr("Not Yours")})
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.ListVersions(c)
var resp map[string]interface{}
@@ -192,7 +194,8 @@ func TestListAgentVersionsHandler_CanvasNotFound(t *testing.T) {
c.Set("user_id", "user-1")
c.Params = gin.Params{{Key: "canvas_id", Value: "non-existent"}}
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.ListVersions(c)
var resp map[string]interface{}
@@ -244,7 +247,8 @@ func TestGetAgentVersionHandler_Success(t *testing.T) {
},
})
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.GetVersion(c)
if w.Code != http.StatusOK {
@@ -291,7 +295,8 @@ func TestGetAgentVersionHandler_VersionNotFound(t *testing.T) {
Title: sptr("Test Agent"),
})
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.GetVersion(c)
var resp map[string]interface{}
@@ -658,7 +663,8 @@ func TestAgentChatCompletions_RequiresAgentID(t *testing.T) {
c.Set("user", &entity.User{ID: "u1"})
c.Set("user_id", "u1")
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.AgentChatCompletions(c)
if w.Code != http.StatusOK {
@@ -686,7 +692,8 @@ func TestAgentChatCompletions_OpenAICompat_EmptyMessages(t *testing.T) {
c.Set("user", &entity.User{ID: "u1"})
c.Set("user_id", "u1")
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.AgentChatCompletions(c)
var resp map[string]interface{}
@@ -998,7 +1005,8 @@ func TestAgentChatCompletions_OpenAICompat_NonStreamReturnsChoices(t *testing.T)
c.Set("user", &entity.User{ID: "u1"})
c.Set("user_id", "u1")
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.AgentChatCompletions(c)
var resp map[string]interface{}
@@ -1035,7 +1043,8 @@ func TestRerunAgent_RequiresAllFields(t *testing.T) {
c.Set("user", &entity.User{ID: "u1"})
c.Set("user_id", "u1")
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.RerunAgent(c)
var resp map[string]interface{}
@@ -1068,7 +1077,8 @@ func TestRerunAgent_AcceptsCompleteRequest(t *testing.T) {
c.Set("user_id", "u1")
stub := &stubDocService{accessible: true}
h := NewAgentHandler(service.NewAgentService(), nil).
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithDocumentService(stub)
h.RerunAgent(c)
@@ -1089,7 +1099,8 @@ func TestPromptsReturnsHardcodedFields(t *testing.T) {
c.Set("user", &entity.User{ID: "u1"})
c.Set("user_id", "u1")
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.Prompts(c)
var resp map[string]interface{}
@@ -1126,7 +1137,8 @@ func TestGetAgentWebhookLogsReturnsEmptyPoll(t *testing.T) {
c.Set("user_id", "u1")
c.Params = gin.Params{{Key: "canvas_id", Value: "c1"}}
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
h.GetAgentWebhookLogs(c)
var resp map[string]interface{}
@@ -1169,7 +1181,8 @@ func TestRerunAgent_RejectsInaccessibleDocument(t *testing.T) {
// round 5), so the deny-all stub injects cleanly without standing
// up the real DocumentService (DB, storage, ...).
stub := &stubDocService{accessible: false}
h := NewAgentHandler(service.NewAgentService(), nil).
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil).
WithDocumentService(stub)
h.RerunAgent(c)
@@ -1201,7 +1214,8 @@ func TestRerunAgent_NoDocumentServiceFailsClosed(t *testing.T) {
c.Set("user", &entity.User{ID: "u1"})
c.Set("user_id", "u1")
h := NewAgentHandler(service.NewAgentService(), nil)
ctx := t.Context()
h := NewAgentHandler(ctx, service.NewAgentService(), nil)
// Note: no WithDocumentService call → documentService is nil.
// Production wiring (cmd/server_main.go) always calls
// WithDocumentService; a nil here means the handler was

View File

@@ -520,7 +520,7 @@ func (h *AgentHandler) runWebhookDetached(
zap.String("canvas", cv.ID),
zap.Error(err))
if isTest {
appendWebhookTrace(cv.ID, startTs, canvas.RunEvent{Type: "error", SessionID: sessionID, Data: mustJSON(map[string]any{"message": err.Error()})})
appendWebhookTrace(ctx, cv.ID, startTs, canvas.RunEvent{Type: "error", SessionID: sessionID, Data: mustJSON(map[string]any{"message": err.Error()})})
}
return
}
@@ -529,7 +529,7 @@ func (h *AgentHandler) runWebhookDetached(
ev.SessionID = sessionID
}
if isTest {
appendWebhookTrace(cv.ID, startTs, ev)
appendWebhookTrace(ctx, cv.ID, startTs, ev)
}
}
}
@@ -553,8 +553,8 @@ func (h *AgentHandler) runWebhookSync(
events, err := h.loader.RunAgentWithWebhook(ctx, cv.UserID, cv.ID, payload)
if err != nil {
if isTest {
appendWebhookTrace(cv.ID, startTs, canvas.RunEvent{Type: "error", SessionID: sessionID, Data: mustJSON(map[string]any{"message": err.Error()})})
appendWebhookTrace(cv.ID, startTs, canvas.RunEvent{Type: "finished", SessionID: sessionID, Data: mustJSON(map[string]any{"success": false})})
appendWebhookTrace(ctx, cv.ID, startTs, canvas.RunEvent{Type: "error", SessionID: sessionID, Data: mustJSON(map[string]any{"message": err.Error()})})
appendWebhookTrace(ctx, cv.ID, startTs, canvas.RunEvent{Type: "finished", SessionID: sessionID, Data: mustJSON(map[string]any{"success": false})})
}
return webhookSyncResult{status: http.StatusBadRequest, body: gin.H{
"code": 400,
@@ -571,7 +571,7 @@ func (h *AgentHandler) runWebhookSync(
ev.SessionID = sessionID
}
if isTest {
appendWebhookTrace(cv.ID, startTs, ev)
appendWebhookTrace(ctx, cv.ID, startTs, ev)
}
switch ev.Type {
case "message":
@@ -602,7 +602,7 @@ func (h *AgentHandler) runWebhookSync(
}
final := strings.Join(contents, "")
if isTest {
appendWebhookTrace(cv.ID, startTs, canvas.RunEvent{Type: "finished", SessionID: sessionID, Data: mustJSON(map[string]any{"success": true})})
appendWebhookTrace(ctx, cv.ID, startTs, canvas.RunEvent{Type: "finished", SessionID: sessionID, Data: mustJSON(map[string]any{"success": true})})
}
return webhookSyncResult{status: status, body: gin.H{
"message": final,
@@ -629,14 +629,14 @@ func mustJSON(v any) string {
// The trace key is `webhook-trace-<agent_id>-logs` with a 600 s TTL.
// Each event is recorded as {"ts": <float>, "event": <type>, ...}.
// Tests use miniredis to verify the key shape.
func appendWebhookTrace(agentID string, startTs time.Time, ev canvas.RunEvent) {
func appendWebhookTrace(ctx context.Context, agentID string, startTs time.Time, ev canvas.RunEvent) {
rdb := rediscli.Get()
if rdb == nil {
return
}
key := fmt.Sprintf("webhook-trace-%s-logs", agentID)
raw, _ := rdb.Get(key)
raw, _ := rdb.Get(ctx, key)
obj := map[string]any{}
if raw != "" {
_ = json.Unmarshal([]byte(raw), &obj)
@@ -677,5 +677,5 @@ func appendWebhookTrace(agentID string, startTs time.Time, ev canvas.RunEvent) {
common.Warn("webhook trace marshal failed", zap.Error(err))
return
}
rdb.SetObj(key, string(encoded), 600*time.Second)
rdb.SetObj(ctx, key, string(encoded), 600*time.Second)
}

View File

@@ -109,7 +109,8 @@ func (h *SystemHandler) GetStatus(c *gin.Context) {
return
}
status, err := h.systemService.GetStatus()
ctx := c.Request.Context()
status, err := h.systemService.GetStatus(ctx)
if err != nil {
jsonInternalError(c, err)
return

View File

@@ -99,7 +99,7 @@ func (h *UserHandler) Register(c *gin.Context) {
return
}
secretKey, err := server.GetSecretKey(redis.Get())
secretKey, err := server.GetSecretKey(ctx, redis.Get())
if err != nil {
common.ResponseWithCodeData(c, common.CodeServerError, false, err.Error())
return
@@ -169,7 +169,7 @@ func (h *UserHandler) Login(c *gin.Context) {
operationLog.UserID = user.ID
// Sign the access_token using itsdangerous (compatible with Python)
secretKey, err := server.GetSecretKey(redis.Get())
secretKey, err := server.GetSecretKey(ctx, redis.Get())
if err != nil {
errMessage := fmt.Sprintf("Failed to get secret key: %s", err.Error())
common.ResponseWithCodeData(c, common.CodeServerError, false, errMessage)
@@ -256,7 +256,7 @@ func (h *UserHandler) LoginByEmail(c *gin.Context) {
}
operationLog.UserID = user.ID
secretKey, err := server.GetSecretKey(redis.Get())
secretKey, err := server.GetSecretKey(ctx, redis.Get())
if err != nil {
errorMessage := fmt.Sprintf("Failed to get secret key: %s", err.Error())
common.ResponseWithCodeData(c, common.CodeServerError, false, errorMessage)
@@ -713,7 +713,7 @@ func (h *UserHandler) ForgotResetPassword(c *gin.Context) {
return
}
secretKey, err := server.GetSecretKey(redis.Get())
secretKey, err := server.GetSecretKey(ctx, redis.Get())
if err != nil {
common.ResponseWithCodeData(c, common.CodeServerError, false, fmt.Sprintf("Failed to get secret key: %s", err.Error()))
return