diff options
Diffstat (limited to 'internal/store/oauth_tokens_test.go')
| -rw-r--r-- | internal/store/oauth_tokens_test.go | 121 |
1 files changed, 121 insertions, 0 deletions
diff --git a/internal/store/oauth_tokens_test.go b/internal/store/oauth_tokens_test.go new file mode 100644 index 0000000..20c4a79 --- /dev/null +++ b/internal/store/oauth_tokens_test.go @@ -0,0 +1,121 @@ +package store + +import ( + "database/sql" + "path/filepath" + "testing" + "time" + + _ "github.com/mattn/go-sqlite3" + "golang.org/x/oauth2" +) + +func newOAuthTokensTestStore(t *testing.T) *Store { + t.Helper() + dbPath := filepath.Join(t.TempDir(), "test.db") + db, err := sql.Open("sqlite3", dbPath) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { db.Close() }) + if _, err := db.Exec(` + CREATE TABLE oauth_tokens ( + source TEXT PRIMARY KEY, + access_token TEXT NOT NULL, + refresh_token TEXT NOT NULL, + token_type TEXT NOT NULL DEFAULT 'Bearer', + expiry DATETIME, + updated_at DATETIME DEFAULT CURRENT_TIMESTAMP + ) + `); err != nil { + t.Fatal(err) + } + return &Store{db: db} +} + +func TestGetOAuthToken_NotConnected_ReturnsNil(t *testing.T) { + s := newOAuthTokensTestStore(t) + + tok, err := s.GetOAuthToken("google_tasks") + if err != nil { + t.Fatalf("GetOAuthToken: %v", err) + } + if tok != nil { + t.Errorf("tok = %+v, want nil", tok) + } +} + +func TestSaveOAuthToken_ThenGet_RoundTrips(t *testing.T) { + s := newOAuthTokensTestStore(t) + + expiry := time.Now().Add(time.Hour).Truncate(time.Second) + in := &oauth2.Token{ + AccessToken: "access-1", + RefreshToken: "refresh-1", + TokenType: "Bearer", + Expiry: expiry, + } + if err := s.SaveOAuthToken("google_tasks", in); err != nil { + t.Fatalf("SaveOAuthToken: %v", err) + } + + got, err := s.GetOAuthToken("google_tasks") + if err != nil { + t.Fatalf("GetOAuthToken: %v", err) + } + if got == nil { + t.Fatal("got nil, want a token") + } + if got.AccessToken != "access-1" || got.RefreshToken != "refresh-1" || got.TokenType != "Bearer" { + t.Errorf("got = %+v, want access-1/refresh-1/Bearer", got) + } + if !got.Expiry.Equal(expiry) { + t.Errorf("Expiry = %v, want %v", got.Expiry, expiry) + } +} + +func TestSaveOAuthToken_Upserts(t *testing.T) { + s := newOAuthTokensTestStore(t) + + if err := s.SaveOAuthToken("google_tasks", &oauth2.Token{AccessToken: "old", RefreshToken: "refresh-1"}); err != nil { + t.Fatal(err) + } + if err := s.SaveOAuthToken("google_tasks", &oauth2.Token{AccessToken: "new", RefreshToken: "refresh-2"}); err != nil { + t.Fatal(err) + } + + got, err := s.GetOAuthToken("google_tasks") + if err != nil { + t.Fatal(err) + } + if got.AccessToken != "new" || got.RefreshToken != "refresh-2" { + t.Errorf("got = %+v, want access token updated to new/refresh-2", got) + } + + var count int + if err := s.db.QueryRow(`SELECT COUNT(*) FROM oauth_tokens`).Scan(&count); err != nil { + t.Fatal(err) + } + if count != 1 { + t.Errorf("row count = %d, want 1 (upsert, not insert)", count) + } +} + +func TestDeleteOAuthToken_RemovesRow(t *testing.T) { + s := newOAuthTokensTestStore(t) + + if err := s.SaveOAuthToken("google_tasks", &oauth2.Token{AccessToken: "a", RefreshToken: "r"}); err != nil { + t.Fatal(err) + } + if err := s.DeleteOAuthToken("google_tasks"); err != nil { + t.Fatalf("DeleteOAuthToken: %v", err) + } + + got, err := s.GetOAuthToken("google_tasks") + if err != nil { + t.Fatal(err) + } + if got != nil { + t.Errorf("got = %+v, want nil after delete", got) + } +} |
