|
5 | 5 | "net/http" |
6 | 6 | "net/http/httptest" |
7 | 7 | "strings" |
| 8 | + "sync" |
8 | 9 | "testing" |
9 | 10 | "time" |
10 | 11 |
|
@@ -201,6 +202,61 @@ func TestRESTEndpoints(t *testing.T) { |
201 | 202 | } |
202 | 203 | } |
203 | 204 |
|
| 205 | +func TestRESTCarriesServeToken(t *testing.T) { |
| 206 | + // Every REST request must carry the per-instance serve token, mirroring |
| 207 | + // odek serve's requireServeToken (cookie or X-Odek-Ws-Token header). |
| 208 | + var seen []string |
| 209 | + var mu sync.Mutex |
| 210 | + record := func(w http.ResponseWriter, r *http.Request) { |
| 211 | + mu.Lock() |
| 212 | + seen = append(seen, r.URL.Path+":"+r.Header.Get("X-Odek-Ws-Token")) |
| 213 | + mu.Unlock() |
| 214 | + } |
| 215 | + mux := http.NewServeMux() |
| 216 | + mux.Handle("/ws", ws.Handler(func(c *ws.Conn) { _, _ = c.Write(nil) })) |
| 217 | + mux.HandleFunc("/api/sessions", func(w http.ResponseWriter, r *http.Request) { |
| 218 | + record(w, r) |
| 219 | + json.NewEncoder(w).Encode([]Session{}) |
| 220 | + }) |
| 221 | + mux.HandleFunc("/api/models", func(w http.ResponseWriter, r *http.Request) { |
| 222 | + record(w, r) |
| 223 | + json.NewEncoder(w).Encode([]ModelInfo{}) |
| 224 | + }) |
| 225 | + mux.HandleFunc("/api/resources", func(w http.ResponseWriter, r *http.Request) { |
| 226 | + record(w, r) |
| 227 | + json.NewEncoder(w).Encode([]Resource{}) |
| 228 | + }) |
| 229 | + mux.HandleFunc("/api/cancel", func(w http.ResponseWriter, r *http.Request) { |
| 230 | + record(w, r) |
| 231 | + w.WriteHeader(http.StatusNoContent) |
| 232 | + }) |
| 233 | + cl, _ := newTestServer(t, mux) |
| 234 | + |
| 235 | + if _, err := cl.Sessions(); err != nil { |
| 236 | + t.Fatalf("Sessions: %v", err) |
| 237 | + } |
| 238 | + if _, err := cl.Models(); err != nil { |
| 239 | + t.Fatalf("Models: %v", err) |
| 240 | + } |
| 241 | + if _, err := cl.Resources("x", 1); err != nil { |
| 242 | + t.Fatalf("Resources: %v", err) |
| 243 | + } |
| 244 | + if err := cl.Cancel("s1", "tok"); err != nil { |
| 245 | + t.Fatalf("Cancel: %v", err) |
| 246 | + } |
| 247 | + |
| 248 | + mu.Lock() |
| 249 | + defer mu.Unlock() |
| 250 | + if len(seen) != 4 { |
| 251 | + t.Fatalf("expected 4 API calls, got %v", seen) |
| 252 | + } |
| 253 | + for _, s := range seen { |
| 254 | + if !strings.HasSuffix(s, ":test-token") { |
| 255 | + t.Errorf("request missing serve token header: %q", s) |
| 256 | + } |
| 257 | + } |
| 258 | +} |
| 259 | + |
204 | 260 | func TestRESTErrorStatuses(t *testing.T) { |
205 | 261 | mux := http.NewServeMux() |
206 | 262 | mux.Handle("/ws", ws.Handler(func(c *ws.Conn) {})) |
|
0 commit comments