package middleware import ( "context" "net/http" "net/http/httptest" "testing" "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/syntax" atp "tangled.org/pdewey.com/atp" ) func TestGetDID_NoContext(t *testing.T) { r := httptest.NewRequest("GET", "/", nil) did, ok := GetDID(r.Context()) if ok || did != "" { t.Fatal("expected no DID in empty context") } } func TestGetSessionID_NoContext(t *testing.T) { r := httptest.NewRequest("GET", "/", nil) sid, ok := GetSessionID(r.Context()) if ok || sid != "" { t.Fatal("expected no session ID in empty context") } } func TestCookieAuth_NoCookies(t *testing.T) { store := oauth.NewMemStore() app, _ := atp.NewOAuthApp(atp.OAuthConfig{ RedirectURI: "http://127.0.0.1:9999/cb", Scopes: []string{"atproto"}, Store: store, }) var gotDID string handler := CookieAuth(CookieAuthConfig{OAuthApp: app})( http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { did, _ := GetDID(r.Context()) gotDID = did }), ) r := httptest.NewRequest("GET", "/", nil) w := httptest.NewRecorder() handler.ServeHTTP(w, r) if gotDID != "" { t.Fatalf("expected empty DID without cookies, got %q", gotDID) } } func TestCookieAuth_ValidSession(t *testing.T) { store := oauth.NewMemStore() did, _ := syntax.ParseDID("did:plc:test123") store.SaveSession(nil, oauth.ClientSessionData{ AccountDID: did, SessionID: "sess-1", }) app, _ := atp.NewOAuthApp(atp.OAuthConfig{ RedirectURI: "http://127.0.0.1:9999/cb", Scopes: []string{"atproto"}, Store: store, }) var gotDID, gotSID string handler := CookieAuth(CookieAuthConfig{OAuthApp: app})( http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { gotDID, _ = GetDID(r.Context()) gotSID, _ = GetSessionID(r.Context()) }), ) r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: "account_did", Value: "did:plc:test123"}) r.AddCookie(&http.Cookie{Name: "session_id", Value: "sess-1"}) w := httptest.NewRecorder() handler.ServeHTTP(w, r) if gotDID != "did:plc:test123" { t.Fatalf("got DID %q, want did:plc:test123", gotDID) } if gotSID != "sess-1" { t.Fatalf("got session ID %q, want sess-1", gotSID) } } func TestCookieAuth_InvalidDID(t *testing.T) { store := oauth.NewMemStore() app, _ := atp.NewOAuthApp(atp.OAuthConfig{ RedirectURI: "http://127.0.0.1:9999/cb", Scopes: []string{"atproto"}, Store: store, }) var gotDID string handler := CookieAuth(CookieAuthConfig{OAuthApp: app})( http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { gotDID, _ = GetDID(r.Context()) }), ) r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: "account_did", Value: "not-a-did"}) r.AddCookie(&http.Cookie{Name: "session_id", Value: "sess-1"}) w := httptest.NewRecorder() handler.ServeHTTP(w, r) if gotDID != "" { t.Fatalf("expected empty DID for invalid cookie, got %q", gotDID) } } func TestCookieAuth_OnAuth(t *testing.T) { store := oauth.NewMemStore() did, _ := syntax.ParseDID("did:plc:test123") store.SaveSession(nil, oauth.ClientSessionData{ AccountDID: did, SessionID: "sess-1", }) app, _ := atp.NewOAuthApp(atp.OAuthConfig{ RedirectURI: "http://127.0.0.1:9999/cb", Scopes: []string{"atproto"}, Store: store, }) var calledWith string handler := CookieAuth(CookieAuthConfig{ OAuthApp: app, OnAuth: func(did string) { calledWith = did }, })(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {})) r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: "account_did", Value: "did:plc:test123"}) r.AddCookie(&http.Cookie{Name: "session_id", Value: "sess-1"}) handler.ServeHTTP(httptest.NewRecorder(), r) if calledWith != "did:plc:test123" { t.Fatalf("OnAuth called with %q, want did:plc:test123", calledWith) } } func TestCookieAuth_CustomCookieNames(t *testing.T) { store := oauth.NewMemStore() did, _ := syntax.ParseDID("did:plc:test123") store.SaveSession(nil, oauth.ClientSessionData{ AccountDID: did, SessionID: "sess-1", }) app, _ := atp.NewOAuthApp(atp.OAuthConfig{ RedirectURI: "http://127.0.0.1:9999/cb", Scopes: []string{"atproto"}, Store: store, }) var gotDID string handler := CookieAuth(CookieAuthConfig{ OAuthApp: app, DIDCookieName: "my_did", SessCookieName: "my_sess", })(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { gotDID, _ = GetDID(r.Context()) })) r := httptest.NewRequest("GET", "/", nil) r.AddCookie(&http.Cookie{Name: "my_did", Value: "did:plc:test123"}) r.AddCookie(&http.Cookie{Name: "my_sess", Value: "sess-1"}) handler.ServeHTTP(httptest.NewRecorder(), r) if gotDID != "did:plc:test123" { t.Fatalf("got DID %q with custom cookie names", gotDID) } } func TestClientMetadataHandler(t *testing.T) { app, _ := atp.NewOAuthApp(atp.OAuthConfig{ RedirectURI: "http://127.0.0.1:9999/cb", Scopes: []string{"atproto"}, AppName: "TestApp", }) handler := ClientMetadataHandler(app) r := httptest.NewRequest("GET", "/client-metadata.json", nil) w := httptest.NewRecorder() handler.ServeHTTP(w, r) if w.Code != http.StatusOK { t.Fatalf("expected 200, got %d", w.Code) } ct := w.Header().Get("Content-Type") if ct != "application/json" { t.Fatalf("expected application/json, got %q", ct) } body := w.Body.String() if !containsStr(body, "client_id") { t.Fatalf("response missing client_id: %s", body) } } func containsStr(s, substr string) bool { for i := 0; i <= len(s)-len(substr); i++ { if s[i:i+len(substr)] == substr { return true } } return false } func TestContextWithAuth(t *testing.T) { ctx := ContextWithAuth(context.Background(), "did:plc:alice", "sess-123") did, ok := GetDID(ctx) if !ok || did != "did:plc:alice" { t.Fatalf("expected did:plc:alice, got %q (ok=%v)", did, ok) } sid, ok := GetSessionID(ctx) if !ok || sid != "sess-123" { t.Fatalf("expected sess-123, got %q (ok=%v)", sid, ok) } } func TestContextWithAuth_EmptyDID(t *testing.T) { ctx := ContextWithAuth(context.Background(), "", "sess-123") _, ok := GetDID(ctx) if ok { t.Fatal("expected GetDID to return false for empty DID") } sid, ok := GetSessionID(ctx) if !ok || sid != "sess-123" { t.Fatalf("expected sess-123, got %q (ok=%v)", sid, ok) } }