|
5 | 5 | "encoding/json" |
6 | 6 | "net/http" |
7 | 7 | "net/http/httptest" |
| 8 | + "os" |
8 | 9 | "path/filepath" |
9 | 10 | "sync" |
10 | 11 | "testing" |
@@ -277,6 +278,70 @@ func TestAgent_ReportStatus(t *testing.T) { |
277 | 278 | assert.Equal(t, 12345, receivedReq.Workers[0].PID) |
278 | 279 | } |
279 | 280 |
|
| 281 | +func TestAgent_ReportStatus_ReadsConnectionFilesEveryTime(t *testing.T) { |
| 282 | + tmpDir := t.TempDir() |
| 283 | + configDir := filepath.Join(tmpDir, "config") |
| 284 | + stateDir := filepath.Join(tmpDir, "state") |
| 285 | + connectionsDir := filepath.Join(tmpDir, "connections") |
| 286 | + require.NoError(t, os.MkdirAll(connectionsDir, 0755)) |
| 287 | + |
| 288 | + var mu sync.Mutex |
| 289 | + var receivedReqs []api.AgentStatusRequest |
| 290 | + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 291 | + var req api.AgentStatusRequest |
| 292 | + _ = json.NewDecoder(r.Body).Decode(&req) |
| 293 | + mu.Lock() |
| 294 | + receivedReqs = append(receivedReqs, req) |
| 295 | + mu.Unlock() |
| 296 | + w.Header().Set("Content-Type", "application/json") |
| 297 | + _ = json.NewEncoder(w).Encode(api.SuccessResponse{Success: true}) |
| 298 | + })) |
| 299 | + defer server.Close() |
| 300 | + |
| 301 | + configMgr := config.NewManager(configDir, stateDir) |
| 302 | + err := configMgr.SaveGPUs([]config.GPUConfig{ |
| 303 | + {GPUID: "GPU-0", GPUIndex: 0, Vendor: "nvidia", Model: "RTX 4090", VRAMMb: 24576}, |
| 304 | + }) |
| 305 | + require.NoError(t, err) |
| 306 | + err = configMgr.SaveWorkers([]config.WorkerConfig{ |
| 307 | + {WorkerID: "worker_1", GPUIDs: []string{"GPU-0"}, ListenPort: 9001, Enabled: true, Status: "running", PID: 12345}, |
| 308 | + }) |
| 309 | + require.NoError(t, err) |
| 310 | + err = os.WriteFile(filepath.Join(connectionsDir, "worker_1.txt"), []byte("10.1.1.1,1234,11\n"), 0644) |
| 311 | + require.NoError(t, err) |
| 312 | + |
| 313 | + client := api.NewClient( |
| 314 | + api.WithBaseURL(server.URL), |
| 315 | + api.WithAgentSecret("gpugo_secret123"), |
| 316 | + ) |
| 317 | + agent := &Agent{ |
| 318 | + client: client, |
| 319 | + config: configMgr, |
| 320 | + ctx: context.Background(), |
| 321 | + agentID: "agent_test123", |
| 322 | + paths: platform.DefaultPaths().WithConfigDir(configDir), |
| 323 | + connectionsDir: connectionsDir, |
| 324 | + } |
| 325 | + |
| 326 | + err = agent.reportStatus() |
| 327 | + require.NoError(t, err) |
| 328 | + |
| 329 | + err = os.WriteFile(filepath.Join(connectionsDir, "worker_1.txt"), []byte("10.2.2.2,5678,22\n"), 0644) |
| 330 | + require.NoError(t, err) |
| 331 | + err = agent.reportStatus() |
| 332 | + require.NoError(t, err) |
| 333 | + |
| 334 | + mu.Lock() |
| 335 | + defer mu.Unlock() |
| 336 | + require.Len(t, receivedReqs, 2) |
| 337 | + require.Len(t, receivedReqs[0].Workers, 1) |
| 338 | + require.Len(t, receivedReqs[0].Workers[0].Connections, 1) |
| 339 | + assert.Equal(t, "10.1.1.1", receivedReqs[0].Workers[0].Connections[0].ClientIP) |
| 340 | + require.Len(t, receivedReqs[1].Workers, 1) |
| 341 | + require.Len(t, receivedReqs[1].Workers[0].Connections, 1) |
| 342 | + assert.Equal(t, "10.2.2.2", receivedReqs[1].Workers[0].Connections[0].ClientIP) |
| 343 | +} |
| 344 | + |
280 | 345 | func TestAgent_HandleHeartbeatResponse(t *testing.T) { |
281 | 346 | tmpDir := t.TempDir() |
282 | 347 | configDir := filepath.Join(tmpDir, "config") |
@@ -492,6 +557,44 @@ func TestAgent_WithHypervisor(t *testing.T) { |
492 | 557 | agent.Stop() |
493 | 558 | } |
494 | 559 |
|
| 560 | +func TestAgent_ConvertToWorkerInfos_IncludesConnectionInfoPath(t *testing.T) { |
| 561 | + tmpDir := t.TempDir() |
| 562 | + configDir := filepath.Join(tmpDir, "config") |
| 563 | + stateDir := filepath.Join(tmpDir, "state") |
| 564 | + |
| 565 | + configMgr := config.NewManager(configDir, stateDir) |
| 566 | + cfg := &config.Config{ |
| 567 | + ConfigVersion: 1, |
| 568 | + AgentID: "agent_test123", |
| 569 | + AgentSecret: "gpugo_secret123", |
| 570 | + ServerURL: "http://localhost", |
| 571 | + License: api.License{ |
| 572 | + Plain: "test|pro|9999999999", |
| 573 | + Encrypted: "enc", |
| 574 | + }, |
| 575 | + } |
| 576 | + err := configMgr.SaveConfig(cfg) |
| 577 | + require.NoError(t, err) |
| 578 | + |
| 579 | + agent := NewAgent(api.NewClient(), configMgr) |
| 580 | + agent.workerBinaryPath = "/bin/true" |
| 581 | + agent.connectionsDir = filepath.Join(tmpDir, "connections") |
| 582 | + |
| 583 | + infos, err := agent.convertToWorkerInfos([]api.WorkerConfig{ |
| 584 | + { |
| 585 | + WorkerID: "worker_1", |
| 586 | + GPUIDs: []string{"gpu-0"}, |
| 587 | + ListenPort: 9001, |
| 588 | + Enabled: true, |
| 589 | + }, |
| 590 | + }) |
| 591 | + require.NoError(t, err) |
| 592 | + require.Len(t, infos, 1) |
| 593 | + require.NotNil(t, infos[0].WorkerRunningInfo) |
| 594 | + |
| 595 | + assert.Equal(t, agent.connectionsDir, infos[0].WorkerRunningInfo.Env[EnvConnectionInfoPath]) |
| 596 | +} |
| 597 | + |
495 | 598 | func TestAgent_LicenseParsing(t *testing.T) { |
496 | 599 | tmpDir := t.TempDir() |
497 | 600 | configDir := filepath.Join(tmpDir, "config") |
|
0 commit comments