store_test.go

  1package atproto
  2
  3import (
  4	"context"
  5	"database/sql"
  6	"testing"
  7	"time"
  8
  9	"github.com/bluesky-social/indigo/atproto/auth/oauth"
 10	"github.com/bluesky-social/indigo/atproto/syntax"
 11	"github.com/stretchr/testify/assert"
 12	_ "github.com/mattn/go-sqlite3"
 13)
 14
 15func makeStore(t *testing.T) (*Store, func()) {
 16	t.Helper()
 17	db, err := sql.Open("sqlite3", "file:test.db?cache=shared&mode=memory&_fk=false")
 18	if err != nil {
 19		t.Fatalf("failed to open db: %v", err)
 20	}
 21	store, err := NewOauthStore(db)
 22	if err != nil {
 23		t.Fatalf("failed to create store: %v", err)
 24	}
 25	return store, func() { _ = db.Close() }
 26}
 27
 28func testDID() syntax.DID {
 29	did, _ := syntax.ParseDID("did:plc:1234567890abcdefghij")
 30	return did
 31}
 32
 33func testSessionData(did syntax.DID, sessionID string) oauth.ClientSessionData {
 34	return oauth.ClientSessionData{
 35		AccountDID:                   did,
 36		SessionID:                    sessionID,
 37		HostURL:                      "https://example.com",
 38		AuthServerURL:                "https://auth.example.com",
 39		AuthServerTokenEndpoint:      "https://auth.example.com/token",
 40		AuthServerRevocationEndpoint: "https://auth.example.com/revoke",
 41		Scopes:                       []string{"atproto"},
 42		AccessToken:                  "access-token-1",
 43		RefreshToken:                 "refresh-token-1",
 44		DPoPAuthServerNonce:         "nonce-auth-1",
 45		DPoPHostNonce:                "nonce-host-1",
 46		DPoPPrivateKeyMultibase:      "zDnaaBxv5C8R7R2Nn9Q3eV4F1G6HbJkLm0PsTtUuVwXyZ",
 47	}
 48}
 49
 50func testAuthRequestData(state string) oauth.AuthRequestData {
 51	return oauth.AuthRequestData{
 52		State:                   state,
 53		AuthServerURL:           "https://auth.example.com",
 54		Scopes:                  []string{"atproto"},
 55		PKCEVerifier:            "pkce-verifier-1",
 56		RequestURI:              "urn:ietf:params:oauth:request_uri:abc123",
 57		AuthServerTokenEndpoint: "https://auth.example.com/token",
 58		DPoPAuthServerNonce:     "nonce-auth-1",
 59		DPoPPrivateKeyMultibase: "zDnaaBxv5C8R7R2Nn9Q3eV4F1G6HbJkLm0PsTtUuVwXyZ",
 60	}
 61}
 62
 63func TestSaveAndGetSession(t *testing.T) {
 64	store, cleanup := makeStore(t)
 65	defer cleanup()
 66	ctx := context.Background()
 67	did := testDID()
 68	sess := testSessionData(did, "session-1")
 69
 70	assert.NoError(t, store.SaveSession(ctx, sess))
 71
 72	got, err := store.GetSession(ctx, did, "session-1")
 73	assert.NoError(t, err)
 74	assert.Equal(t, sess.AccountDID.String(), got.AccountDID.String())
 75	assert.Equal(t, sess.SessionID, got.SessionID)
 76	assert.Equal(t, sess.AccessToken, got.AccessToken)
 77	assert.Equal(t, sess.RefreshToken, got.RefreshToken)
 78	assert.Equal(t, sess.Scopes, got.Scopes)
 79	assert.Equal(t, sess.DPoPPrivateKeyMultibase, got.DPoPPrivateKeyMultibase)
 80}
 81
 82func TestSaveSessionUpsert(t *testing.T) {
 83	store, cleanup := makeStore(t)
 84	defer cleanup()
 85	ctx := context.Background()
 86	did := testDID()
 87
 88	sess1 := testSessionData(did, "session-1")
 89	sess1.AccessToken = "access-token-1"
 90	sess1.RefreshToken = "refresh-token-1"
 91	assert.NoError(t, store.SaveSession(ctx, sess1))
 92
 93	sess2 := testSessionData(did, "session-1")
 94	sess2.AccessToken = "access-token-2"
 95	sess2.RefreshToken = "refresh-token-2"
 96	assert.NoError(t, store.SaveSession(ctx, sess2))
 97
 98	var count int
 99	err := store.db.QueryRow("SELECT COUNT(*) FROM oauth_session WHERE did=? AND session_id=?", did.String(), "session-1").Scan(&count)
100	assert.NoError(t, err)
101	assert.Equal(t, 1, count, "upsert should not create duplicate rows")
102
103	got, err := store.GetSession(ctx, did, "session-1")
104	assert.NoError(t, err)
105	assert.Equal(t, "access-token-2", got.AccessToken)
106	assert.Equal(t, "refresh-token-2", got.RefreshToken)
107
108	var mtime string
109	err = store.db.QueryRow("SELECT mtime FROM oauth_session WHERE did=? AND session_id=?", did.String(), "session-1").Scan(&mtime)
110	assert.NoError(t, err)
111	assert.NotEmpty(t, mtime)
112}
113
114func TestGetSessionNotFound(t *testing.T) {
115	store, cleanup := makeStore(t)
116	defer cleanup()
117	ctx := context.Background()
118	did := testDID()
119
120	_, err := store.GetSession(ctx, did, "nonexistent-session")
121	assert.Error(t, err)
122	assert.ErrorIs(t, err, sql.ErrNoRows)
123}
124
125func TestDeleteSession(t *testing.T) {
126	store, cleanup := makeStore(t)
127	defer cleanup()
128	ctx := context.Background()
129	did := testDID()
130	sess := testSessionData(did, "session-1")
131	assert.NoError(t, store.SaveSession(ctx, sess))
132
133	assert.NoError(t, store.DeleteSession(ctx, did, "session-1"))
134
135	_, err := store.GetSession(ctx, did, "session-1")
136	assert.Error(t, err)
137}
138
139func TestMultipleSessionsForSameDID(t *testing.T) {
140	store, cleanup := makeStore(t)
141	defer cleanup()
142	ctx := context.Background()
143	did := testDID()
144
145	sess1 := testSessionData(did, "session-1")
146	sess1.AccessToken = "token-1"
147	sess2 := testSessionData(did, "session-2")
148	sess2.AccessToken = "token-2"
149
150	assert.NoError(t, store.SaveSession(ctx, sess1))
151	assert.NoError(t, store.SaveSession(ctx, sess2))
152
153	got1, err := store.GetSession(ctx, did, "session-1")
154	assert.NoError(t, err)
155	assert.Equal(t, "token-1", got1.AccessToken)
156
157	got2, err := store.GetSession(ctx, did, "session-2")
158	assert.NoError(t, err)
159	assert.Equal(t, "token-2", got2.AccessToken)
160
161	assert.NoError(t, store.DeleteSession(ctx, did, "session-1"))
162
163	_, err = store.GetSession(ctx, did, "session-1")
164	assert.Error(t, err)
165
166	got2After, err := store.GetSession(ctx, did, "session-2")
167	assert.NoError(t, err)
168	assert.Equal(t, "token-2", got2After.AccessToken)
169}
170
171func TestSaveAndGetAuthRequestInfo(t *testing.T) {
172	store, cleanup := makeStore(t)
173	defer cleanup()
174	ctx := context.Background()
175	info := testAuthRequestData("state-123")
176
177	assert.NoError(t, store.SaveAuthRequestInfo(ctx, info))
178
179	got, err := store.GetAuthRequestInfo(ctx, "state-123")
180	assert.NoError(t, err)
181	assert.Equal(t, info.State, got.State)
182	assert.Equal(t, info.AuthServerURL, got.AuthServerURL)
183	assert.Equal(t, info.PKCEVerifier, got.PKCEVerifier)
184	assert.Equal(t, info.RequestURI, got.RequestURI)
185	assert.Equal(t, info.DPoPAuthServerNonce, got.DPoPAuthServerNonce)
186	assert.Equal(t, info.DPoPPrivateKeyMultibase, got.DPoPPrivateKeyMultibase)
187	assert.Equal(t, info.Scopes, got.Scopes)
188}
189
190func TestSaveAuthRequestInfoCreateOnly(t *testing.T) {
191	store, cleanup := makeStore(t)
192	defer cleanup()
193	ctx := context.Background()
194	info := testAuthRequestData("state-123")
195	assert.NoError(t, store.SaveAuthRequestInfo(ctx, info))
196
197	dup := testAuthRequestData("state-123")
198	dup.PKCEVerifier = "different-verifier"
199	err := store.SaveAuthRequestInfo(ctx, dup)
200	assert.Error(t, err)
201
202	got, err := store.GetAuthRequestInfo(ctx, "state-123")
203	assert.NoError(t, err)
204	assert.Equal(t, "pkce-verifier-1", got.PKCEVerifier)
205}
206
207func TestGetAuthRequestInfoNotFound(t *testing.T) {
208	store, cleanup := makeStore(t)
209	defer cleanup()
210	ctx := context.Background()
211
212	_, err := store.GetAuthRequestInfo(ctx, "nonexistent-state")
213	assert.Error(t, err)
214	assert.ErrorIs(t, err, sql.ErrNoRows)
215}
216
217func TestDeleteAuthRequestInfo(t *testing.T) {
218	store, cleanup := makeStore(t)
219	defer cleanup()
220	ctx := context.Background()
221	info := testAuthRequestData("state-123")
222	assert.NoError(t, store.SaveAuthRequestInfo(ctx, info))
223
224	assert.NoError(t, store.DeleteAuthRequestInfo(ctx, "state-123"))
225
226	_, err := store.GetAuthRequestInfo(ctx, "state-123")
227	assert.Error(t, err)
228}
229
230func TestMigrateIdempotent(t *testing.T) {
231	db, err := sql.Open("sqlite3", "file:test.db?cache=shared&mode=memory&_fk=false")
232	assert.NoError(t, err)
233	defer func() { _ = db.Close() }()
234
235	assert.NoError(t, migrate(db))
236	assert.NoError(t, migrate(db))
237
238	store := &Store{db: db}
239	ctx := context.Background()
240	did := testDID()
241	sess := testSessionData(did, "session-1")
242	assert.NoError(t, store.SaveSession(ctx, sess))
243	got, err := store.GetSession(ctx, did, "session-1")
244	assert.NoError(t, err)
245	assert.Equal(t, sess.AccessToken, got.AccessToken)
246}
247
248func TestGCExpiredSessions(t *testing.T) {
249	store, cleanup := makeStore(t)
250	defer cleanup()
251	ctx := context.Background()
252	did := testDID()
253
254	sess := testSessionData(did, "old-session")
255	assert.NoError(t, store.SaveSession(ctx, sess))
256	_, err := store.db.Exec("UPDATE oauth_session SET mtime=? WHERE did=? AND session_id=?",
257		time.Now().UTC().Add(-7*30*24*time.Hour).Format(SQLITE_TIMEFMT),
258		did.String(), "old-session")
259	assert.NoError(t, err)
260
261	freshSess := testSessionData(did, "fresh-session")
262	assert.NoError(t, store.SaveSession(ctx, freshSess))
263
264	var count int
265	err = store.db.QueryRow("SELECT COUNT(*) FROM oauth_session WHERE session_id=?", "old-session").Scan(&count)
266	assert.NoError(t, err)
267	assert.Equal(t, 0, count, "expired session should be deleted by GC trigger")
268
269	err = store.db.QueryRow("SELECT COUNT(*) FROM oauth_session WHERE session_id=?", "fresh-session").Scan(&count)
270	assert.NoError(t, err)
271	assert.Equal(t, 1, count, "fresh session should not be deleted")
272}
273
274func TestGCExpiredAuthRequests(t *testing.T) {
275	store, cleanup := makeStore(t)
276	defer cleanup()
277	ctx := context.Background()
278
279	info := testAuthRequestData("old-state")
280	assert.NoError(t, store.SaveAuthRequestInfo(ctx, info))
281	_, err := store.db.Exec("UPDATE oauth_request SET mtime=? WHERE state=?",
282		time.Now().UTC().Add(-7*30*24*time.Hour).Format(SQLITE_TIMEFMT),
283		"old-state")
284	assert.NoError(t, err)
285
286	freshInfo := testAuthRequestData("fresh-state")
287	assert.NoError(t, store.SaveAuthRequestInfo(ctx, freshInfo))
288
289	var count int
290	err = store.db.QueryRow("SELECT COUNT(*) FROM oauth_request WHERE state=?", "old-state").Scan(&count)
291	assert.NoError(t, err)
292	assert.Equal(t, 0, count, "expired auth request should be deleted by GC trigger")
293
294	err = store.db.QueryRow("SELECT COUNT(*) FROM oauth_request WHERE state=?", "fresh-state").Scan(&count)
295	assert.NoError(t, err)
296	assert.Equal(t, 1, count, "fresh auth request should not be deleted")
297}
298