205 lines
5.6 KiB
Go
205 lines
5.6 KiB
Go
package ai_client_test
|
|
|
|
import (
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"sync/atomic"
|
|
"testing"
|
|
|
|
"thuanle.me/claw-email-bridge/internal/ai_client"
|
|
"thuanle.me/claw-email-bridge/internal/config"
|
|
"thuanle.me/claw-email-bridge/internal/database"
|
|
)
|
|
|
|
func TestBuildDispatchMetadata_ForwardsContext(t *testing.T) {
|
|
task := &database.Task{
|
|
TaskUUID: "uuid-ctx-test",
|
|
DispatchContext: `{"rule":"whitelist","source":"env"}`,
|
|
}
|
|
metadata := ai_client.ExportBuildDispatchMetadata(task)
|
|
|
|
if metadata["task_uuid"] != "uuid-ctx-test" {
|
|
t.Errorf("expected task_uuid=uuid-ctx-test, got %q", metadata["task_uuid"])
|
|
}
|
|
if metadata["rule"] != "whitelist" {
|
|
t.Errorf("expected rule=whitelist, got %q", metadata["rule"])
|
|
}
|
|
if metadata["source"] != "env" {
|
|
t.Errorf("expected source=env, got %q", metadata["source"])
|
|
}
|
|
}
|
|
|
|
func TestBuildDispatchMetadata_ProtectsTaskUUID(t *testing.T) {
|
|
task := &database.Task{
|
|
TaskUUID: "original-uuid",
|
|
DispatchContext: `{"task_uuid":"malicious-override","extra":"data"}`,
|
|
}
|
|
metadata := ai_client.ExportBuildDispatchMetadata(task)
|
|
|
|
if metadata["task_uuid"] != "original-uuid" {
|
|
t.Errorf("task_uuid should not be overwritten, got %q", metadata["task_uuid"])
|
|
}
|
|
if metadata["extra"] != "data" {
|
|
t.Errorf("expected extra=data, got %q", metadata["extra"])
|
|
}
|
|
}
|
|
|
|
func TestBuildDispatchMetadata_NoContext(t *testing.T) {
|
|
task := &database.Task{
|
|
TaskUUID: "uuid-no-ctx",
|
|
}
|
|
metadata := ai_client.ExportBuildDispatchMetadata(task)
|
|
|
|
if len(metadata) != 1 {
|
|
t.Errorf("expected 1 key, got %d", len(metadata))
|
|
}
|
|
if metadata["task_uuid"] != "uuid-no-ctx" {
|
|
t.Errorf("expected task_uuid=uuid-no-ctx, got %q", metadata["task_uuid"])
|
|
}
|
|
}
|
|
|
|
// Test matrix #2: OpenClaw fail 2 lần, lần 3 thành công → COMPLETED, attempt_openclaw = 3.
|
|
func TestDispatch_RetryThenSuccess(t *testing.T) {
|
|
tdb := database.NewTestDB(t)
|
|
|
|
var callCount atomic.Int32
|
|
mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
n := callCount.Add(1)
|
|
if n <= 2 {
|
|
w.WriteHeader(http.StatusInternalServerError)
|
|
w.Write([]byte("temporary error"))
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusOK)
|
|
w.Write([]byte(`{"status":"ok"}`))
|
|
}))
|
|
defer mockServer.Close()
|
|
|
|
cfg := &config.Config{
|
|
OpenClawURL: mockServer.URL,
|
|
ListenAddr: ":9999",
|
|
}
|
|
|
|
task := &database.Task{
|
|
TaskUUID: "uuid-retry-test",
|
|
ThreadID: "thread-1",
|
|
MessageID: "msg-retry@test.com",
|
|
Sender: "user@test.com",
|
|
Subject: "Retry test",
|
|
BodyPlain: "Hello",
|
|
Status: database.StatusReceived,
|
|
}
|
|
if err := database.CreateTask(tdb.DB, task); err != nil {
|
|
t.Fatalf("create task: %v", err)
|
|
}
|
|
|
|
dispatcher := ai_client.NewDispatcher(cfg, tdb.DB)
|
|
dispatcher.Dispatch(task) // Synchronous call for testing.
|
|
|
|
// Verify 3 attempts were made.
|
|
if callCount.Load() != 3 {
|
|
t.Errorf("expected 3 API calls, got %d", callCount.Load())
|
|
}
|
|
|
|
// Reload task from DB.
|
|
updated, err := database.FindByTaskUUID(tdb.DB, "uuid-retry-test")
|
|
if err != nil {
|
|
t.Fatalf("find task: %v", err)
|
|
}
|
|
if updated.Status != database.StatusAIProcessing {
|
|
t.Errorf("expected AI_PROCESSING (waiting for callback), got %s", updated.Status)
|
|
}
|
|
if updated.AttemptOpenClaw != 3 {
|
|
t.Errorf("expected attempt_openclaw = 3, got %d", updated.AttemptOpenClaw)
|
|
}
|
|
}
|
|
|
|
// OpenClaw fail all 3 retries → FAILED with last_error.
|
|
func TestDispatch_AllRetriesFailed(t *testing.T) {
|
|
tdb := database.NewTestDB(t)
|
|
|
|
mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusInternalServerError)
|
|
w.Write([]byte("server down"))
|
|
}))
|
|
defer mockServer.Close()
|
|
|
|
cfg := &config.Config{
|
|
OpenClawURL: mockServer.URL,
|
|
ListenAddr: ":9999",
|
|
}
|
|
|
|
task := &database.Task{
|
|
TaskUUID: "uuid-fail-all",
|
|
ThreadID: "thread-1",
|
|
MessageID: "msg-fail@test.com",
|
|
Sender: "user@test.com",
|
|
Subject: "Fail test",
|
|
BodyPlain: "Hello",
|
|
Status: database.StatusReceived,
|
|
}
|
|
if err := database.CreateTask(tdb.DB, task); err != nil {
|
|
t.Fatalf("create task: %v", err)
|
|
}
|
|
|
|
dispatcher := ai_client.NewDispatcher(cfg, tdb.DB)
|
|
dispatcher.Dispatch(task)
|
|
|
|
updated, err := database.FindByTaskUUID(tdb.DB, "uuid-fail-all")
|
|
if err != nil {
|
|
t.Fatalf("find task: %v", err)
|
|
}
|
|
if updated.Status != database.StatusFailed {
|
|
t.Errorf("expected FAILED, got %s", updated.Status)
|
|
}
|
|
if updated.LastError == nil {
|
|
t.Error("expected last_error to be set")
|
|
}
|
|
if updated.AttemptOpenClaw != 3 {
|
|
t.Errorf("expected attempt_openclaw = 3, got %d", updated.AttemptOpenClaw)
|
|
}
|
|
}
|
|
|
|
func TestDispatch_UsesCallbackBaseURL(t *testing.T) {
|
|
tdb := database.NewTestDB(t)
|
|
|
|
var callbackURL string
|
|
mockServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
var payload map[string]any
|
|
if err := json.NewDecoder(r.Body).Decode(&payload); err != nil {
|
|
t.Fatalf("decode payload: %v", err)
|
|
}
|
|
val, _ := payload["callback_url"].(string)
|
|
callbackURL = val
|
|
w.WriteHeader(http.StatusOK)
|
|
}))
|
|
defer mockServer.Close()
|
|
|
|
cfg := &config.Config{
|
|
OpenClawURL: mockServer.URL,
|
|
ListenAddr: ":9999",
|
|
CallbackBaseURL: "https://bridge.example.com/",
|
|
}
|
|
|
|
task := &database.Task{
|
|
TaskUUID: "uuid-callback-url",
|
|
ThreadID: "thread-1",
|
|
MessageID: "msg-callback@test.com",
|
|
Sender: "user@test.com",
|
|
Subject: "Callback test",
|
|
BodyPlain: "Hello",
|
|
Status: database.StatusReceived,
|
|
}
|
|
if err := database.CreateTask(tdb.DB, task); err != nil {
|
|
t.Fatalf("create task: %v", err)
|
|
}
|
|
|
|
dispatcher := ai_client.NewDispatcher(cfg, tdb.DB)
|
|
dispatcher.Dispatch(task)
|
|
|
|
if callbackURL != "https://bridge.example.com/callback" {
|
|
t.Fatalf("expected normalized callback URL, got %q", callbackURL)
|
|
}
|
|
}
|