Skip to content

Commit 1295572

Browse files
committed
Fix cloud model refresh fallbacks
Refresh OpenCSG models on startup and recover from empty or unauthorized cloud model lists so Chat keeps showing available OpenCSG and provider models.
1 parent 83190c9 commit 1295572

9 files changed

Lines changed: 254 additions & 13 deletions

internal/cloud/opencsg.go

Lines changed: 45 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -197,29 +197,67 @@ func (s *Service) cachedModels() ([]api.ModelInfo, bool) {
197197

198198
func (s *Service) refresh(ctx context.Context) ([]api.ModelInfo, error) {
199199
baseURL := s.BaseURL()
200+
token := s.currentAccessToken()
201+
models, limits, err := fetchCloudModels(ctx, s.client, baseURL, token)
202+
if err != nil && token != "" && isUnauthorizedStatus(err) {
203+
models, limits, err = fetchCloudModels(ctx, s.client, baseURL, "")
204+
}
205+
if err != nil {
206+
return nil, err
207+
}
208+
209+
s.mu.Lock()
210+
s.cached = cloneModels(models)
211+
s.limits = limits
212+
s.cachedAt = time.Now()
213+
s.mu.Unlock()
214+
215+
return models, nil
216+
}
217+
218+
type cloudModelListStatusError struct {
219+
statusCode int
220+
body string
221+
}
222+
223+
func (e cloudModelListStatusError) Error() string {
224+
return fmt.Sprintf("cloud model list returned %d: %s", e.statusCode, e.body)
225+
}
226+
227+
func isUnauthorizedStatus(err error) bool {
228+
if statusErr, ok := err.(cloudModelListStatusError); ok {
229+
return statusErr.statusCode == http.StatusUnauthorized || statusErr.statusCode == http.StatusForbidden
230+
}
231+
return false
232+
}
233+
234+
func fetchCloudModels(ctx context.Context, client *http.Client, baseURL, token string) ([]api.ModelInfo, map[string]ModelTokenLimits, error) {
200235
req, err := http.NewRequestWithContext(ctx, http.MethodGet, baseURL+"/v1/models?page="+cloudModelListPage+"&per="+cloudModelListPer, nil)
201236
if err != nil {
202-
return nil, fmt.Errorf("creating cloud model request: %w", err)
237+
return nil, nil, fmt.Errorf("creating cloud model request: %w", err)
203238
}
204239
req.Header.Set("Accept", "application/json")
205-
if token := s.currentAccessToken(); token != "" {
240+
if token != "" {
206241
req.Header.Set("Authorization", "Bearer "+token)
207242
}
208243

209-
resp, err := s.client.Do(req)
244+
resp, err := client.Do(req)
210245
if err != nil {
211-
return nil, fmt.Errorf("fetching cloud models: %w", err)
246+
return nil, nil, fmt.Errorf("fetching cloud models: %w", err)
212247
}
213248
defer resp.Body.Close()
214249

215250
if resp.StatusCode != http.StatusOK {
216251
body, _ := io.ReadAll(io.LimitReader(resp.Body, 4096))
217-
return nil, fmt.Errorf("cloud model list returned %d: %s", resp.StatusCode, strings.TrimSpace(string(body)))
252+
return nil, nil, cloudModelListStatusError{
253+
statusCode: resp.StatusCode,
254+
body: strings.TrimSpace(string(body)),
255+
}
218256
}
219257

220258
var payload modelListResponse
221259
if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil {
222-
return nil, fmt.Errorf("decoding cloud model list: %w", err)
260+
return nil, nil, fmt.Errorf("decoding cloud model list: %w", err)
223261
}
224262

225263
models := make([]api.ModelInfo, 0, len(payload.Data))
@@ -233,13 +271,7 @@ func (s *Service) refresh(ctx context.Context) ([]api.ModelInfo, error) {
233271
limits[strings.TrimSpace(info.Model)] = modelTokenLimitsFromRemote(item)
234272
}
235273

236-
s.mu.Lock()
237-
s.cached = cloneModels(models)
238-
s.limits = limits
239-
s.cachedAt = time.Now()
240-
s.mu.Unlock()
241-
242-
return models, nil
274+
return models, limits, nil
243275
}
244276

245277
func (s *Service) currentAccessToken() string {

internal/cloud/opencsg_test.go

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -312,6 +312,42 @@ func TestRefreshChatModelsOmitsAuthorizationWithoutAccessToken(t *testing.T) {
312312
}
313313
}
314314

315+
func TestRefreshChatModelsFallsBackToPublicListWhenAccessTokenUnauthorized(t *testing.T) {
316+
requests := 0
317+
apiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
318+
requests++
319+
if got := r.Header.Get("Authorization"); got != "" {
320+
http.Error(w, "expired token", http.StatusUnauthorized)
321+
return
322+
}
323+
w.Header().Set("Content-Type", "application/json")
324+
_ = json.NewEncoder(w).Encode(map[string]any{
325+
"object": "list",
326+
"data": []map[string]any{
327+
{
328+
"id": "public/model",
329+
"task": "text-generation",
330+
},
331+
},
332+
})
333+
}))
334+
defer apiServer.Close()
335+
336+
svc := NewService(apiServer.URL)
337+
svc.SetAccessToken("expired-token")
338+
339+
models, err := svc.RefreshChatModels(context.Background())
340+
if err != nil {
341+
t.Fatalf("RefreshChatModels returned error: %v", err)
342+
}
343+
if requests != 2 {
344+
t.Fatalf("requests = %d, want authenticated request plus public fallback", requests)
345+
}
346+
if len(models) != 1 || models[0].Model != "public/model" {
347+
t.Fatalf("models = %#v, want public/model", models)
348+
}
349+
}
350+
315351
func TestRefreshChatModelsBypassesCache(t *testing.T) {
316352
requests := 0
317353
apiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {

internal/server/cloud_models.go

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@ package server
22

33
import (
44
"context"
5+
"log"
56
"net/http"
67
"sort"
78
"strings"
@@ -12,6 +13,7 @@ import (
1213
)
1314

1415
const forcedCloudModelRefreshInterval = 30 * time.Second
16+
const startupCloudModelRefreshTimeout = 20 * time.Second
1517

1618
func requestWantsModelRefresh(r *http.Request) bool {
1719
value := strings.TrimSpace(strings.ToLower(r.URL.Query().Get("refresh")))
@@ -23,6 +25,21 @@ func requestWantsModelRefresh(r *http.Request) bool {
2325
}
2426
}
2527

28+
func (s *Server) refreshCloudModelsOnStartup(parent context.Context) {
29+
if s == nil || s.cloud == nil {
30+
return
31+
}
32+
ctx, cancel := context.WithTimeout(parent, startupCloudModelRefreshTimeout)
33+
defer cancel()
34+
35+
models, err := s.refreshCloudChatModels(ctx)
36+
if err != nil {
37+
log.Printf("startup cloud model refresh failed: %v", err)
38+
return
39+
}
40+
log.Printf("startup cloud model refresh complete: %d models", len(models))
41+
}
42+
2643
func (s *Server) listAvailableModelsWithRefresh(ctx context.Context, refreshCloud bool) ([]api.ModelInfo, error) {
2744
localModels, err := s.listLocalModelInfos()
2845
if err != nil {
@@ -56,6 +73,8 @@ func (s *Server) listAvailableModelsWithRefresh(ctx context.Context, refreshClou
5673
seen[modelID] = struct{}{}
5774
out = append(out, item)
5875
}
76+
} else {
77+
log.Printf("cloud model list unavailable: %v", err)
5978
}
6079

6180
for _, item := range s.listSelectedThirdPartyProviderModels(ctx) {
@@ -88,6 +107,9 @@ func (s *Server) listCloudModels(ctx context.Context, refresh bool) ([]api.Model
88107
return s.withConfiguredCloudProvider(models), err
89108
}
90109
models, err := s.cloud.ListChatModels(ctx)
110+
if err == nil && len(models) == 0 {
111+
models, err = s.refreshCloudChatModels(ctx)
112+
}
91113
return s.withConfiguredCloudProvider(models), err
92114
}
93115

internal/server/cloud_models_test.go

Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -75,6 +75,47 @@ func TestHandleTagsWithoutTokenIncludesAndRefreshesCloudModels(t *testing.T) {
7575
}
7676
}
7777

78+
func TestHandleTagsRefreshesCloudModelsWhenCacheIsEmpty(t *testing.T) {
79+
requests := 0
80+
apiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
81+
requests++
82+
w.Header().Set("Content-Type", "application/json")
83+
data := []map[string]any{}
84+
if requests > 1 {
85+
data = append(data, map[string]any{
86+
"id": "fallback/model",
87+
"task": "text-generation",
88+
})
89+
}
90+
_ = json.NewEncoder(w).Encode(map[string]any{
91+
"object": "list",
92+
"data": data,
93+
})
94+
}))
95+
defer apiServer.Close()
96+
97+
s := newTestServer(t)
98+
s.cloud = cloud.NewService(apiServer.URL)
99+
100+
req := httptest.NewRequest(http.MethodGet, "/api/tags", nil)
101+
w := httptest.NewRecorder()
102+
s.handleTags(w, req)
103+
if w.Code != http.StatusOK {
104+
t.Fatalf("status = %d, want %d", w.Code, http.StatusOK)
105+
}
106+
107+
var resp api.TagsResponse
108+
if err := json.NewDecoder(w.Body).Decode(&resp); err != nil {
109+
t.Fatalf("decode tags response: %v", err)
110+
}
111+
if requests != 2 {
112+
t.Fatalf("requests = %d, want initial empty fetch plus fallback refresh", requests)
113+
}
114+
if len(resp.Models) != 1 || resp.Models[0].Model != "fallback/model" {
115+
t.Fatalf("models = %#v, want fallback/model", resp.Models)
116+
}
117+
}
118+
78119
func TestHandleTagsSendsLoginTokenToCloudGateway(t *testing.T) {
79120
apiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
80121
if got := r.Header.Get("Authorization"); got != "Bearer access-token" {

internal/server/handlers_audio.go

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -132,6 +132,7 @@ func (s *Server) streamAudioTranscription(w http.ResponseWriter, r *http.Request
132132
})
133133
if err != nil {
134134
log.Printf("MODEL %s: ASR stream transcription failed: %v", modelID, err)
135+
s.closeASREngine(modelID)
135136
writeSSE(w, map[string]interface{}{
136137
"error": err.Error(),
137138
"done": true,

internal/server/handlers_audio_test.go

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@ package server
22

33
import (
44
"bytes"
5+
"context"
6+
"errors"
57
"io"
68
"mime/multipart"
79
"net/http"
@@ -12,6 +14,7 @@ import (
1214
"testing"
1315

1416
"github.com/opencsgs/csghub-lite/internal/config"
17+
"github.com/opencsgs/csghub-lite/pkg/api"
1518
)
1619

1720
type zeroReader struct{}
@@ -21,6 +24,27 @@ func (zeroReader) Read(p []byte) (int, error) {
2124
return len(p), nil
2225
}
2326

27+
type failingStreamASREngine struct {
28+
closed bool
29+
}
30+
31+
func (e *failingStreamASREngine) Transcribe(context.Context, api.OpenAIAudioTranscriptionRequest) (*api.OpenAIAudioTranscriptionResponse, error) {
32+
return nil, errors.New("not used")
33+
}
34+
35+
func (e *failingStreamASREngine) TranscribeStream(context.Context, api.OpenAIAudioTranscriptionRequest, func(api.OpenAIAudioTranscriptionResponse) error) error {
36+
return errors.New("worker exited")
37+
}
38+
39+
func (e *failingStreamASREngine) Close() error {
40+
e.closed = true
41+
return nil
42+
}
43+
44+
func (e *failingStreamASREngine) ModelName() string {
45+
return "test-asr"
46+
}
47+
2448
func TestHandleOpenAIAudioTranscriptionsUsesLiteTempDir(t *testing.T) {
2549
missingTempDir := filepath.Join(t.TempDir(), "missing-temp")
2650
t.Setenv("TMPDIR", missingTempDir)
@@ -100,3 +124,28 @@ func TestHandleOpenAIAudioTranscriptionsParsesFieldsFromStreamedMultipart(t *tes
100124
t.Fatalf("expected streamed multipart fields to be parsed, got status=%d body=%s", w.Code, w.Body.String())
101125
}
102126
}
127+
128+
func TestStreamAudioTranscriptionClosesFailedWorker(t *testing.T) {
129+
modelID := "AIWizards/Fun-ASR-Nano-2512"
130+
engine := &failingStreamASREngine{}
131+
s := New(&config.Config{
132+
ModelDir: config.ModelDirForStorage(t.TempDir()),
133+
DatasetDir: config.DatasetDirForStorage(t.TempDir()),
134+
}, "test")
135+
s.asrEngines[modelID] = &managedASREngine{engine: engine}
136+
137+
req := httptest.NewRequest(http.MethodPost, "/v1/audio/transcriptions", nil)
138+
w := httptest.NewRecorder()
139+
140+
s.streamAudioTranscription(w, req, modelID, engine, api.OpenAIAudioTranscriptionRequest{})
141+
142+
if !engine.closed {
143+
t.Fatal("expected failed ASR stream worker to be closed")
144+
}
145+
if _, ok := s.asrEngines[modelID]; ok {
146+
t.Fatal("expected failed ASR stream worker to be removed from cache")
147+
}
148+
if !strings.Contains(w.Body.String(), "worker exited") {
149+
t.Fatalf("expected SSE error response, got %q", w.Body.String())
150+
}
151+
}

internal/server/handlers_cloud.go

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -192,6 +192,9 @@ func (s *Server) getChatEngine(ctx context.Context, modelID, source string, numC
192192

193193
models, cloudErr := s.listCloudModels(ctx, false)
194194
if cloudErr != nil {
195+
if providerSource := s.thirdPartyProviderSourceForModel(ctx, modelID); providerSource != "" {
196+
return newThirdPartyProviderEngine(providerSource, modelID)
197+
}
195198
return nil, err
196199
}
197200
if modelInfoListContains(models, modelID) {
@@ -203,6 +206,9 @@ func (s *Server) getChatEngine(ctx context.Context, modelID, source string, numC
203206
}
204207
models, cloudErr = s.cloud.RefreshChatModels(ctx)
205208
if cloudErr != nil {
209+
if providerSource := s.thirdPartyProviderSourceForModel(ctx, modelID); providerSource != "" {
210+
return newThirdPartyProviderEngine(providerSource, modelID)
211+
}
206212
return nil, err
207213
}
208214
if modelInfoListContains(models, modelID) {

internal/server/handlers_providers_test.go

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1041,6 +1041,59 @@ func TestGetChatEnginePrefersThirdPartyWhenCloudLoginMissing(t *testing.T) {
10411041
}
10421042
}
10431043

1044+
func TestGetChatEngineFallsBackToThirdPartyWhenCloudListFails(t *testing.T) {
1045+
home := t.TempDir()
1046+
t.Setenv("HOME", home)
1047+
t.Setenv("USERPROFILE", home)
1048+
config.ResetProviders()
1049+
config.ResetProviderModelAllowlist()
1050+
t.Cleanup(config.ResetProviders)
1051+
t.Cleanup(config.ResetProviderModelAllowlist)
1052+
1053+
providerServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
1054+
switch r.URL.Path {
1055+
case "/v1/chat/completions":
1056+
w.Header().Set("Content-Type", "application/json")
1057+
_, _ = fmt.Fprint(w, `{"choices":[{"message":{"role":"assistant","content":"provider ok"}}]}`)
1058+
default:
1059+
http.NotFound(w, r)
1060+
}
1061+
}))
1062+
defer providerServer.Close()
1063+
cloudServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
1064+
http.Error(w, "cloud unavailable", http.StatusInternalServerError)
1065+
}))
1066+
defer cloudServer.Close()
1067+
1068+
s := newTestServer(t)
1069+
s.cfg.OpenCSGAPIKey = "cloud-key"
1070+
s.cloud = cloud.NewService(cloudServer.URL)
1071+
if err := config.SaveProviders([]config.ThirdPartyProvider{{
1072+
ID: "provider1",
1073+
Name: "OpenAI",
1074+
BaseURL: providerServer.URL + "/v1",
1075+
APIKey: "secret",
1076+
Enabled: true,
1077+
}}); err != nil {
1078+
t.Fatalf("save providers: %v", err)
1079+
}
1080+
if err := config.ReplaceProviderModelAllowlist("provider1", []string{"deepseek-v4-pro"}); err != nil {
1081+
t.Fatalf("save provider model allowlist: %v", err)
1082+
}
1083+
1084+
eng, err := s.getChatEngine(context.Background(), "deepseek-v4-pro", "", 0, 0, -1, "", "", "")
1085+
if err != nil {
1086+
t.Fatalf("getChatEngine returned error: %v", err)
1087+
}
1088+
got, err := eng.Chat(context.Background(), nil, inference.DefaultOptions(), nil)
1089+
if err != nil {
1090+
t.Fatalf("chat returned error: %v", err)
1091+
}
1092+
if got != "provider ok" {
1093+
t.Fatalf("chat = %q, want provider ok", got)
1094+
}
1095+
}
1096+
10441097
func TestDisabledProviderExcludedFromTagsAndEngine(t *testing.T) {
10451098
home := t.TempDir()
10461099
t.Setenv("HOME", home)

0 commit comments

Comments
 (0)