package rm import ( "context" "encoding/base64" "io" "net/http" "net/http/httptest" "strings" "testing" "time" "spark-mcp-go/internal/cluster" "spark-mcp-go/internal/httpclient" ) func testClient(srv *httptest.Server, c *cluster.Cluster) *Client { hc := httpclient.New(httpclient.Config{Timeout: 5 * time.Second, MaxResponseBytes: 1 << 20}) return New(hc, c) } func TestListApps(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/ws/v1/cluster/apps" { t.Errorf("path=%q, want /ws/v1/cluster/apps", r.URL.Path) } w.Header().Set("Content-Type", "application/json") w.Write([]byte(`{"apps":{"app":[]}}`)) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL} raw, err := testClient(srv, cl).ListApps(context.Background(), "", "") if err != nil { t.Fatalf("ListApps: %v", err) } if string(raw) != `{"apps":{"app":[]}}` { t.Errorf("body=%q", string(raw)) } } func TestListApps_WithQuery(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { q := r.URL.Query() if q.Get("state") != "RUNNING" { t.Errorf("state=%q, want RUNNING", q.Get("state")) } if q.Get("user") != "yarn" { t.Errorf("user=%q, want yarn", q.Get("user")) } w.Header().Set("Content-Type", "application/json") w.Write([]byte(`{"apps":{"app":[{"id":"app_1"}]}}`)) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL} raw, err := testClient(srv, cl).ListApps(context.Background(), "RUNNING", "yarn") if err != nil { t.Fatalf("ListApps: %v", err) } if !strings.Contains(string(raw), "app_1") { t.Errorf("body=%q", string(raw)) } } func TestGetApp(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/ws/v1/cluster/apps/app_123" { t.Errorf("path=%q", r.URL.Path) } w.Header().Set("Content-Type", "application/json") w.Write([]byte(`{"app":{"id":"app_123","state":"RUNNING"}}`)) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL} raw, err := testClient(srv, cl).GetApp(context.Background(), "app_123") if err != nil { t.Fatalf("GetApp: %v", err) } if !strings.Contains(string(raw), "RUNNING") { t.Errorf("body=%q", string(raw)) } } func TestKillApp(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/ws/v1/cluster/apps/app_123/state" { t.Errorf("path=%q", r.URL.Path) } if r.Method != http.MethodPut { t.Errorf("method=%q, want PUT", r.Method) } body, _ := io.ReadAll(r.Body) if string(body) != `{"state":"KILLED"}` { t.Errorf("body=%q", string(body)) } w.Header().Set("Content-Type", "application/json") w.Write([]byte(`{"state":"KILLED"}`)) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL} raw, err := testClient(srv, cl).KillApp(context.Background(), "app_123") if err != nil { t.Fatalf("KillApp: %v", err) } if string(raw) != `{"state":"KILLED"}` { t.Errorf("body=%q", string(raw)) } } func TestGetLogs_PrimarySuccess(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path == "/ws/v1/cluster/apps/app_123/amContainerLogs" { w.Write([]byte("am logs")) return } t.Errorf("unexpected path %q", r.URL.Path) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL} body, source, err := testClient(srv, cl).GetLogs(context.Background(), "app_123", "container_1") if err != nil { t.Fatalf("GetLogs: %v", err) } if source != "amContainerLogs" { t.Errorf("source=%q, want amContainerLogs", source) } if string(body) != "am logs" { t.Errorf("body=%q", string(body)) } } func TestGetLogs_PrimaryRedirects(t *testing.T) { srvNM := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { auth := r.Header.Get("Authorization") if auth == "" { t.Error("Authorization header missing after redirect") } w.Write([]byte("nm logs")) })) defer srvNM.Close() srvRM := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path == "/ws/v1/cluster/apps/app_123/amContainerLogs" { w.Header().Set("Location", srvNM.URL+"/node/containerlogs/container_1/root") w.WriteHeader(http.StatusTemporaryRedirect) return } t.Errorf("unexpected RM path %q", r.URL.Path) })) defer srvRM.Close() nmHost := srvNM.Listener.Addr().String() cl := &cluster.Cluster{ RMURL: srvRM.URL, AuthType: cluster.AuthBasic, AuthUsername: "u", AuthPassword: "p", URLAllowlist: []string{nmHost}, } body, source, err := testClient(srvRM, cl).GetLogs(context.Background(), "app_123", "container_1") if err != nil { t.Fatalf("GetLogs: %v", err) } if source != "amContainerLogs" { t.Errorf("source=%q, want amContainerLogs", source) } if string(body) != "nm logs" { t.Errorf("body=%q", string(body)) } } func TestGetLogs_FallbackToAggregated(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/ws/v1/cluster/apps/app_123/amContainerLogs": w.WriteHeader(http.StatusInternalServerError) case "/ws/v1/cluster/apps/app_123/aggregated-logs": if r.URL.Query().Get("container") != "container_1" { t.Errorf("container=%q, want container_1", r.URL.Query().Get("container")) } w.Write([]byte("aggregated logs")) default: t.Errorf("unexpected path %q", r.URL.Path) } })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL} body, source, err := testClient(srv, cl).GetLogs(context.Background(), "app_123", "container_1") if err != nil { t.Fatalf("GetLogs: %v", err) } if source != "aggregated-logs" { t.Errorf("source=%q, want aggregated-logs", source) } if string(body) != "aggregated logs" { t.Errorf("body=%q", string(body)) } } func TestApplyAuth_SimpleUserName(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if got := r.URL.Query().Get("user.name"); got != "alice" { t.Errorf("user.name=%q, want alice", got) } w.Write([]byte(`{"apps":{"app":[]}}`)) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL, AuthType: cluster.AuthSimple, AuthUsername: "alice"} _, err := testClient(srv, cl).ListApps(context.Background(), "", "") if err != nil { t.Fatalf("ListApps: %v", err) } } func TestApplyAuth_BasicHeader(t *testing.T) { captured := "" srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { captured = r.Header.Get("Authorization") w.Write([]byte(`{"apps":{"app":[]}}`)) })) defer srv.Close() cl := &cluster.Cluster{RMURL: srv.URL, AuthType: cluster.AuthBasic, AuthUsername: "u", AuthPassword: "p"} _, err := testClient(srv, cl).ListApps(context.Background(), "", "") if err != nil { t.Fatalf("ListApps: %v", err) } want := "Basic " + base64.StdEncoding.EncodeToString([]byte("u:p")) if captured != want { t.Errorf("Authorization=%q, want %q", captured, want) } }