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