Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 124 additions & 0 deletions internal/api/redis_queue_protocol_integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,68 @@ func readRESPArrayOfBulkStrings(r *bufio.Reader) ([][]byte, error) {
return out, nil
}

func readTestRESPPubSubSubscribe(r *bufio.Reader) (string, int, error) {
prefix, errRead := r.ReadByte()
if errRead != nil {
return "", 0, errRead
}
if prefix != '*' {
return "", 0, fmt.Errorf("expected array prefix '*', got %q", prefix)
}
line, errLine := readTestRESPLine(r)
if errLine != nil {
return "", 0, errLine
}
count, errParse := strconv.Atoi(line)
if errParse != nil {
return "", 0, fmt.Errorf("invalid array length %q: %v", line, errParse)
}
if count != 3 {
return "", 0, fmt.Errorf("subscribe ack length = %d, want 3", count)
}
kind, errKind := readTestRESPBulkString(r)
if errKind != nil {
return "", 0, errKind
}
if string(kind) != "subscribe" {
return "", 0, fmt.Errorf("subscribe ack kind = %q", string(kind))
}
channel, errChannel := readTestRESPBulkString(r)
if errChannel != nil {
return "", 0, errChannel
}
prefix, errRead = r.ReadByte()
if errRead != nil {
return "", 0, errRead
}
if prefix != ':' {
return "", 0, fmt.Errorf("expected integer prefix ':', got %q", prefix)
}
line, errLine = readTestRESPLine(r)
if errLine != nil {
return "", 0, errLine
}
subscriptions, errParse := strconv.Atoi(line)
if errParse != nil {
return "", 0, fmt.Errorf("invalid subscription count %q: %v", line, errParse)
}
return string(channel), subscriptions, nil
}

func readTestRESPPubSubMessage(r *bufio.Reader) (string, []byte, error) {
items, errItems := readRESPArrayOfBulkStrings(r)
if errItems != nil {
return "", nil, errItems
}
if len(items) != 3 {
return "", nil, fmt.Errorf("pubsub message length = %d, want 3", len(items))
}
if string(items[0]) != "message" {
return "", nil, fmt.Errorf("pubsub message kind = %q", string(items[0]))
}
return string(items[1]), items[2], nil
}

func TestRedisProtocol_ManagementDisabled_RejectsConnection(t *testing.T) {
t.Setenv("MANAGEMENT_PASSWORD", "")
redisqueue.SetEnabled(false)
Expand Down Expand Up @@ -235,6 +297,68 @@ func TestRedisProtocol_HomeEnabled_DisablesConnection(t *testing.T) {
}
}

func TestRedisProtocol_SUBSCRIBE_UsageSendsSupportRefresh(t *testing.T) {
const managementPassword = "test-management-password"

t.Setenv("MANAGEMENT_PASSWORD", managementPassword)
redisqueue.SetEnabled(false)
t.Cleanup(func() { redisqueue.SetEnabled(false) })

server := newTestServer(t)
if !server.managementRoutesEnabled.Load() {
t.Fatalf("expected managementRoutesEnabled to be true")
}

addr, stop := startRedisMuxListener(t, server)
t.Cleanup(stop)

conn, errDial := net.DialTimeout("tcp", addr, time.Second)
if errDial != nil {
t.Fatalf("failed to dial redis listener: %v", errDial)
}
t.Cleanup(func() { _ = conn.Close() })

reader := bufio.NewReader(conn)
_ = conn.SetDeadline(time.Now().Add(5 * time.Second))

if errWrite := writeTestRESPCommand(conn, "AUTH", managementPassword); errWrite != nil {
t.Fatalf("failed to write AUTH command: %v", errWrite)
}
if msg, errRead := readTestRESPSimpleString(reader); errRead != nil {
t.Fatalf("failed to read AUTH response: %v", errRead)
} else if msg != "OK" {
t.Fatalf("unexpected AUTH response: %q", msg)
}

if errWrite := writeTestRESPCommand(conn, "SUBSCRIBE", "usage"); errWrite != nil {
t.Fatalf("failed to write SUBSCRIBE command: %v", errWrite)
}
channel, subscriptions, errSubscribe := readTestRESPPubSubSubscribe(reader)
if errSubscribe != nil {
t.Fatalf("failed to read subscribe response: %v", errSubscribe)
}
if channel != "usage" || subscriptions != 1 {
t.Fatalf("unexpected subscribe response channel=%q subscriptions=%d", channel, subscriptions)
}

channel, payload, errMessage := readTestRESPPubSubMessage(reader)
if errMessage != nil {
t.Fatalf("failed to read support refresh message: %v", errMessage)
}
if channel != "usage" || string(payload) != `{"support_refresh":true}` {
t.Fatalf("unexpected support refresh message channel=%q payload=%q", channel, string(payload))
}

redisqueue.Enqueue([]byte(`{"id":1}`))
channel, payload, errMessage = readTestRESPPubSubMessage(reader)
if errMessage != nil {
t.Fatalf("failed to read usage message: %v", errMessage)
}
if channel != "usage" || string(payload) != `{"id":1}` {
t.Fatalf("unexpected usage message channel=%q payload=%q", channel, string(payload))
}
}

func TestRedisProtocol_AUTH_And_PopContracts(t *testing.T) {
const managementPassword = "test-management-password"

Expand Down
8 changes: 8 additions & 0 deletions internal/redisqueue/queue.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,9 @@ const (
defaultRetentionSeconds int64 = 60
maxRetentionSeconds int64 = 3600
usageSubscriberBuffer = 256

usageSupportRefreshPayload = `{"support_refresh":true}`
usageRefreshPayload = `{"refresh":true}`
)

type queueItem struct {
Expand Down Expand Up @@ -83,6 +86,10 @@ func SubscribeUsage() (<-chan []byte, func()) {
return global.subscribeUsage()
}

func NotifyUsageRefresh() {
global.publishToSubscribers([]byte(usageRefreshPayload))
}

func (q *queue) clear() {
q.mu.Lock()

Expand Down Expand Up @@ -137,6 +144,7 @@ func (q *queue) publishToSubscribers(payload []byte) bool {

func (q *queue) subscribeUsage() (<-chan []byte, func()) {
subscriber := make(chan []byte, usageSubscriberBuffer)
subscriber <- []byte(usageSupportRefreshPayload)

q.mu.Lock()
if q.subscribers == nil {
Expand Down
23 changes: 23 additions & 0 deletions internal/redisqueue/queue_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,9 @@ func TestEnqueueBroadcastsToUsageSubscribersAndSkipsQueue(t *testing.T) {
second, unsubscribeSecond := SubscribeUsage()
defer unsubscribeSecond()

requireUsageSubscriberPayload(t, first, usageSupportRefreshPayload)
requireUsageSubscriberPayload(t, second, usageSupportRefreshPayload)

Enqueue([]byte("usage-record"))

requireUsageSubscriberPayload(t, first, "usage-record")
Expand All @@ -37,6 +40,8 @@ func TestSetEnabledFalseClosesUsageSubscribers(t *testing.T) {
subscriber, unsubscribe := SubscribeUsage()
defer unsubscribe()

requireUsageSubscriberPayload(t, subscriber, usageSupportRefreshPayload)

SetEnabled(false)

select {
Expand All @@ -50,6 +55,24 @@ func TestSetEnabledFalseClosesUsageSubscribers(t *testing.T) {
})
}

func TestNotifyUsageRefreshBroadcastsOnlyToUsageSubscribers(t *testing.T) {
withEnabledQueue(t, func() {
subscriber, unsubscribe := SubscribeUsage()
defer unsubscribe()

requireUsageSubscriberPayload(t, subscriber, usageSupportRefreshPayload)

NotifyUsageRefresh()
requireUsageSubscriberPayload(t, subscriber, usageRefreshPayload)

unsubscribe()
NotifyUsageRefresh()
if items := PopOldest(1); len(items) != 0 {
t.Fatalf("PopOldest() items = %q, want empty after refresh notification without subscribers", items)
}
})
}

func requireUsageSubscriberPayload(t *testing.T, subscriber <-chan []byte, want string) {
t.Helper()

Expand Down
4 changes: 4 additions & 0 deletions internal/watcher/clients.go
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ import (
"time"

"github.com/router-for-me/CLIProxyAPI/v7/internal/config"
"github.com/router-for-me/CLIProxyAPI/v7/internal/redisqueue"
"github.com/router-for-me/CLIProxyAPI/v7/internal/util"
"github.com/router-for-me/CLIProxyAPI/v7/internal/watcher/diff"
"github.com/router-for-me/CLIProxyAPI/v7/internal/watcher/synthesizer"
Expand Down Expand Up @@ -134,6 +135,7 @@ func (w *Watcher) reloadClients(rescanAuth bool, affectedOAuthProviders []string
}

w.refreshAuthState(forceAuthRefresh)
redisqueue.NotifyUsageRefresh()

log.Infof("full client load complete - %d clients (%d auth files + %d Gemini API keys + %d Vertex API keys + %d Claude API keys + %d Codex keys + %d OpenAI-compat)",
totalNewClients,
Expand Down Expand Up @@ -233,6 +235,7 @@ func (w *Watcher) addOrUpdateClient(path string) {

w.persistAuthAsync(fmt.Sprintf("Sync auth %s", filepath.Base(path)), path)
w.dispatchAuthUpdates(updates)
redisqueue.NotifyUsageRefresh()
}

func (w *Watcher) removeClient(path string) {
Expand All @@ -251,6 +254,7 @@ func (w *Watcher) removeClient(path string) {

w.persistAuthAsync(fmt.Sprintf("Remove auth %s", filepath.Base(path)), path)
w.dispatchAuthUpdates(updates)
redisqueue.NotifyUsageRefresh()
}

func (w *Watcher) computePerPathUpdatesLocked(oldByID, newByID map[string]*coreauth.Auth) []AuthUpdate {
Expand Down
64 changes: 64 additions & 0 deletions internal/watcher/watcher_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ import (

"github.com/fsnotify/fsnotify"
"github.com/router-for-me/CLIProxyAPI/v7/internal/config"
"github.com/router-for-me/CLIProxyAPI/v7/internal/redisqueue"
"github.com/router-for-me/CLIProxyAPI/v7/internal/watcher/diff"
"github.com/router-for-me/CLIProxyAPI/v7/internal/watcher/synthesizer"
sdkAuth "github.com/router-for-me/CLIProxyAPI/v7/sdk/auth"
Expand Down Expand Up @@ -441,6 +442,34 @@ func TestRemoveClientRemovesHash(t *testing.T) {
}
}

func TestAuthFileClientChangesNotifyUsageSubscribersToRefresh(t *testing.T) {
tmpDir := t.TempDir()
authFile := filepath.Join(tmpDir, "sample.json")
if err := os.WriteFile(authFile, []byte(`{"type":"demo","api_key":"k"}`), 0o644); err != nil {
t.Fatalf("failed to create auth file: %v", err)
}

redisqueue.SetEnabled(false)
redisqueue.SetEnabled(true)
t.Cleanup(func() { redisqueue.SetEnabled(false) })

subscriber, unsubscribe := redisqueue.SubscribeUsage()
defer unsubscribe()
requireWatcherUsagePayload(t, subscriber, `{"support_refresh":true}`)

w := &Watcher{
authDir: tmpDir,
lastAuthHashes: make(map[string]string),
}
w.SetConfig(&config.Config{AuthDir: tmpDir})

w.addOrUpdateClient(authFile)
requireWatcherUsagePayload(t, subscriber, `{"refresh":true}`)

w.removeClient(authFile)
requireWatcherUsagePayload(t, subscriber, `{"refresh":true}`)
}

func TestAuthFileEventsDoNotInvokeSnapshotCoreAuths(t *testing.T) {
tmpDir := t.TempDir()
authFile := filepath.Join(tmpDir, "sample.json")
Expand Down Expand Up @@ -699,6 +728,25 @@ func TestReloadClientsHandlesNilConfig(t *testing.T) {
w.reloadClients(true, nil, false)
}

func TestReloadClientsNotifiesUsageSubscribersToRefresh(t *testing.T) {
tmp := t.TempDir()
redisqueue.SetEnabled(false)
redisqueue.SetEnabled(true)
t.Cleanup(func() { redisqueue.SetEnabled(false) })

subscriber, unsubscribe := redisqueue.SubscribeUsage()
defer unsubscribe()
requireWatcherUsagePayload(t, subscriber, `{"support_refresh":true}`)

w := &Watcher{
authDir: tmp,
config: &config.Config{AuthDir: tmp},
}
w.reloadClients(false, nil, false)

requireWatcherUsagePayload(t, subscriber, `{"refresh":true}`)
}

func TestReloadClientsFiltersProvidersWithNilCurrentAuths(t *testing.T) {
tmp := t.TempDir()
w := &Watcher{
Expand All @@ -711,6 +759,22 @@ func TestReloadClientsFiltersProvidersWithNilCurrentAuths(t *testing.T) {
}
}

func requireWatcherUsagePayload(t *testing.T, subscriber <-chan []byte, want string) {
t.Helper()

select {
case got, ok := <-subscriber:
if !ok {
t.Fatalf("subscriber closed before receiving %q", want)
}
if string(got) != want {
t.Fatalf("subscriber payload = %q, want %q", string(got), want)
}
case <-time.After(time.Second):
t.Fatalf("timeout waiting for subscriber payload %q", want)
}
}

func TestSetAuthUpdateQueueNilResetsDispatch(t *testing.T) {
w := &Watcher{}
queue := make(chan AuthUpdate, 1)
Expand Down
Loading