diff --git a/TODO.md b/TODO.md index 85bd198..d60fc65 100644 --- a/TODO.md +++ b/TODO.md @@ -66,5 +66,5 @@ - [ ] Error handling: toast/notification for network failures, auth expiry - [ ] Keyboard shortcuts: `Cmd+K` focus search, `Cmd+R` refresh, `Cmd+L` toggle log viewer - [ ] Window title and app icon (`build/appicon.png`) -- [ ] Production build verification (`wails3 build` → macOS `.app` bundle) +- [ ] Production build verification (`wails build` → macOS `.app` bundle) - [ ] README with build instructions, screenshots, and usage diff --git a/auth_service.go b/auth_service.go index 7547aaa..7266661 100644 --- a/auth_service.go +++ b/auth_service.go @@ -8,7 +8,6 @@ import ( "os/exec" rt "runtime" "strings" - "time" "github.com/bluesky-social/indigo/atproto/auth/oauth" "github.com/bluesky-social/indigo/atproto/identity" @@ -36,27 +35,28 @@ func NewAuthService() *AuthService { // Login initiates OAuth login flow for the given handle func (s *AuthService) Login(handle string) error { ctx := context.Background() + s.codeChan = make(chan string, 1) + s.errChan = make(chan error, 1) - listener, err := net.Listen("tcp", "127.0.0.1:0") + listener, err := net.Listen("tcp", listenerAddress()) if err != nil { return fmt.Errorf("failed to start listener: %w", err) } s.listener = listener - s.port = listener.Addr().(*net.TCPAddr).Port + s.port = oauthCallbackPort - redirectURI := fmt.Sprintf("http://127.0.0.1:%d/callback", s.port) - scopes := []string{"atproto", "transition:generic"} - - config := oauth.NewLocalhostConfig(redirectURI, scopes) - store := oauth.NewMemStore() - s.app = oauth.NewClientApp(&config, store) + store := NewSQLiteOAuthStore() + s.app = newOAuthApp(store) redirectURL, err := s.app.StartAuthFlow(ctx, handle) if err != nil { + closeCallbackServer(nil, s.listener) + s.listener = nil return fmt.Errorf("failed to start auth flow: %w", err) } s.startCallbackServer() + defer s.stopCallbackServer() if err := openBrowser(redirectURL); err != nil { return fmt.Errorf("failed to open browser: %w", err) @@ -107,6 +107,13 @@ func (s *AuthService) startCallbackServer() { }() } +func (s *AuthService) stopCallbackServer() { + closeCallbackServer(s.server, s.listener) + s.server = nil + s.listener = nil + s.port = 0 +} + func (s *AuthService) exchangeCode(ctx context.Context, data string) error { parts := strings.SplitN(data, "|", 3) if len(parts) < 2 { @@ -125,22 +132,18 @@ func (s *AuthService) exchangeCode(ctx context.Context, data string) error { return fmt.Errorf("failed to process callback: %w", err) } - auth := &Auth{ - DID: sessData.AccountDID.String(), - Handle: sessData.AccountDID.String(), - AccessJWT: sessData.AccessToken, - RefreshJWT: sessData.RefreshToken, - PDSURL: sessData.HostURL, - SessionID: sessData.SessionID, - AuthServerURL: sessData.AuthServerURL, - AuthServerTokenEndpoint: sessData.AuthServerTokenEndpoint, - AuthServerRevocationEndpoint: sessData.AuthServerRevocationEndpoint, - DPoPAuthNonce: sessData.DPoPAuthServerNonce, - DPoPHostNonce: sessData.DPoPHostNonce, - DPoPPrivateKey: sessData.DPoPPrivateKeyMultibase, - UpdatedAt: time.Now(), + current, err := GetAuthByDID(sessData.AccountDID.String()) + if err != nil { + return fmt.Errorf("failed to load persisted auth: %w", err) + } + + handle := "" + if current != nil { + handle = current.Handle } + auth := authFromSessionData(sessData, handle) + if err := UpsertAuth(auth); err != nil { return fmt.Errorf("failed to persist auth: %w", err) } @@ -173,6 +176,9 @@ func (s *AuthService) Whoami(force bool) (*Auth, error) { } auth.Handle = ident.Handle.String() + if err := UpsertAuth(auth); err != nil { + return nil, fmt.Errorf("failed to persist resolved handle: %w", err) + } } @@ -202,11 +208,8 @@ func (s *AuthService) RefreshSession() error { return nil // Cannot refresh without session ID } - redirectURI := "http://127.0.0.1/callback" - scopes := []string{"atproto", "transition:generic"} - config := oauth.NewLocalhostConfig(redirectURI, scopes) - store := oauth.NewMemStore() - app := oauth.NewClientApp(&config, store) + store := NewSQLiteOAuthStore() + app := newOAuthApp(store) did, err := syntax.ParseDID(auth.DID) if err != nil { @@ -218,17 +221,12 @@ func (s *AuthService) RefreshSession() error { return fmt.Errorf("failed to resume session: %w", err) } - newAccessToken, err := session.RefreshTokens(context.Background()) - if err != nil { + if _, err := session.RefreshTokens(context.Background()); err != nil { return fmt.Errorf("failed to refresh tokens: %w", err) } - if newAccessToken != "" { - auth.AccessJWT = newAccessToken - auth.UpdatedAt = time.Now() - if err := UpsertAuth(auth); err != nil { - return fmt.Errorf("failed to update refreshed tokens: %w", err) - } + if err := UpsertAuth(authFromSessionData(session.Data, auth.Handle)); err != nil { + return fmt.Errorf("failed to persist refreshed session: %w", err) } return nil diff --git a/database.go b/database.go index bfd8b5b..8d13ebb 100644 --- a/database.go +++ b/database.go @@ -6,6 +6,7 @@ import ( "fmt" "os" "path/filepath" + "strings" "time" _ "modernc.org/sqlite" @@ -193,14 +194,51 @@ func GetAuth() (*Auth, error) { query := `SELECT did, handle, access_jwt, refresh_jwt, pds_url, session_id, auth_server_url, auth_server_token_endpoint, auth_server_revocation_endpoint, dpop_auth_nonce, dpop_host_nonce, dpop_private_key, updated_at - FROM auth LIMIT 1` + FROM auth + ORDER BY updated_at DESC + LIMIT 1` + auth, err := getAuthByQuery(query) + + if err == sql.ErrNoRows { + fmt.Println("no auth record found in database") + return nil, nil + } + if err != nil { + fmt.Printf("failed to load auth: %v\n", err) + return nil, err + } + + fmt.Printf("auth loaded successfully: %s (%s)\n", auth.DID, auth.Handle) + return auth, nil +} + +// GetAuthByDID loads auth for a specific DID. +func GetAuthByDID(did string) (*Auth, error) { + query := `SELECT did, handle, access_jwt, refresh_jwt, pds_url, session_id, + auth_server_url, auth_server_token_endpoint, auth_server_revocation_endpoint, + dpop_auth_nonce, dpop_host_nonce, dpop_private_key, updated_at + FROM auth + WHERE did = ? + LIMIT 1` + + auth, err := getAuthByQuery(query, did) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + return auth, nil +} + +func getAuthByQuery(query string, args ...any) (*Auth, error) { var auth Auth var updatedAt string var sessionID, authServerURL, authServerTokenEndpoint, authServerRevocationEndpoint, dpopAuthNonce, dpopHostNonce, dpopPrivateKey sql.NullString - err := db.QueryRow(query).Scan( + err := db.QueryRow(query, args...).Scan( &auth.DID, &auth.Handle, &auth.AccessJWT, @@ -215,6 +253,9 @@ func GetAuth() (*Auth, error) { &dpopPrivateKey, &updatedAt, ) + if err != nil { + return nil, err + } if sessionID.Valid { auth.SessionID = sessionID.String @@ -238,24 +279,23 @@ func GetAuth() (*Auth, error) { auth.DPoPPrivateKey = dpopPrivateKey.String } - if err == sql.ErrNoRows { - fmt.Println("no auth record found in database") - return nil, nil - } - if err != nil { - fmt.Printf("failed to load auth: %v\n", err) - return nil, err - } - auth.UpdatedAt, _ = time.Parse("2006-01-02 15:04:05", updatedAt) - fmt.Printf("auth loaded successfully: %s (%s)\n", auth.DID, auth.Handle) return &auth, nil } // SearchPosts searches posts using FTS5 func SearchPosts(query string, source string) ([]SearchResult, error) { + query = strings.TrimSpace(query) + if query == "*" { + query = "" + } + fmt.Printf("searching posts: query=%s, source=%s\n", query, source) + if query == "" { + return listRecentPosts(source) + } + sqlQuery := ` SELECT p.uri, p.cid, p.author_did, p.author_handle, p.text, p.created_at, p.like_count, p.repost_count, p.reply_count, p.source, p.indexed_at, @@ -307,6 +347,52 @@ func SearchPosts(query string, source string) ([]SearchResult, error) { return results, rows.Err() } +func listRecentPosts(source string) ([]SearchResult, error) { + rows, err := db.Query(` + SELECT uri, cid, author_did, author_handle, text, created_at, + like_count, repost_count, reply_count, source, indexed_at + FROM posts + WHERE (? = '' OR source = ?) + ORDER BY created_at DESC + LIMIT 25 + `, source, source) + if err != nil { + fmt.Printf("failed to list recent posts: %v\n", err) + return nil, err + } + defer rows.Close() + + var results []SearchResult + for rows.Next() { + var r SearchResult + var createdAt, indexedAt string + + err := rows.Scan( + &r.URI, + &r.CID, + &r.AuthorDID, + &r.AuthorHandle, + &r.Text, + &createdAt, + &r.LikeCount, + &r.RepostCount, + &r.ReplyCount, + &r.Source, + &indexedAt, + ) + if err != nil { + return nil, err + } + + r.CreatedAt, _ = time.Parse("2006-01-02 15:04:05", createdAt) + r.IndexedAt, _ = time.Parse("2006-01-02 15:04:05", indexedAt) + results = append(results, r) + } + + fmt.Printf("browse completed: %d results\n", len(results)) + return results, rows.Err() +} + // CountPosts returns the total number of posts in the database func CountPosts() (int, error) { fmt.Println("counting posts in database") diff --git a/database_test.go b/database_test.go new file mode 100644 index 0000000..f1cfb19 --- /dev/null +++ b/database_test.go @@ -0,0 +1,139 @@ +package main + +import ( + "context" + "path/filepath" + "testing" + "time" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" +) + +func openTestDB(t *testing.T) { + t.Helper() + + dbPath := filepath.Join(t.TempDir(), "test.db") + if err := Open(dbPath); err != nil { + t.Fatalf("Open() error = %v", err) + } + + t.Cleanup(func() { + if err := Close(); err != nil { + t.Fatalf("Close() error = %v", err) + } + }) +} + +func TestSearchPostsBrowseMode(t *testing.T) { + openTestDB(t) + + posts := []*Post{ + { + URI: "at://did:plc:test/app.bsky.feed.post/1", + CID: "cid-1", + AuthorDID: "did:plc:test", + AuthorHandle: "alice.test", + Text: "older saved post", + CreatedAt: time.Date(2026, 3, 14, 12, 0, 0, 0, time.UTC), + Source: "saved", + }, + { + URI: "at://did:plc:test/app.bsky.feed.post/2", + CID: "cid-2", + AuthorDID: "did:plc:test", + AuthorHandle: "alice.test", + Text: "newer liked post", + CreatedAt: time.Date(2026, 3, 15, 12, 0, 0, 0, time.UTC), + Source: "liked", + }, + } + + for _, post := range posts { + if err := InsertPost(post); err != nil { + t.Fatalf("InsertPost() error = %v", err) + } + } + + results, err := SearchPosts("", "") + if err != nil { + t.Fatalf("SearchPosts(empty) error = %v", err) + } + if len(results) != 2 { + t.Fatalf("SearchPosts(empty) len = %d, want 2", len(results)) + } + if results[0].URI != posts[1].URI { + t.Fatalf("SearchPosts(empty) first URI = %q, want %q", results[0].URI, posts[1].URI) + } + + starResults, err := SearchPosts("*", "saved") + if err != nil { + t.Fatalf("SearchPosts(*) error = %v", err) + } + if len(starResults) != 1 { + t.Fatalf("SearchPosts(*) len = %d, want 1", len(starResults)) + } + if starResults[0].Source != "saved" { + t.Fatalf("SearchPosts(*) source = %q, want %q", starResults[0].Source, "saved") + } +} + +func TestSQLiteOAuthStorePersistsSession(t *testing.T) { + openTestDB(t) + + store := NewSQLiteOAuthStore() + did, err := syntax.ParseDID("did:plc:xg2vq45muivyy3xwatcehspu") + if err != nil { + t.Fatalf("ParseDID() error = %v", err) + } + + session := oauth.ClientSessionData{ + AccountDID: did, + SessionID: "session-123", + HostURL: "https://bsky.social", + AuthServerURL: "https://auth.example.com", + AuthServerTokenEndpoint: "https://auth.example.com/token", + AuthServerRevocationEndpoint: "https://auth.example.com/revoke", + Scopes: append([]string(nil), oauthScopes...), + AccessToken: "access-1", + RefreshToken: "refresh-1", + DPoPAuthServerNonce: "auth-nonce", + DPoPHostNonce: "host-nonce", + DPoPPrivateKeyMultibase: "private-key", + } + + if err := store.SaveSession(context.Background(), session); err != nil { + t.Fatalf("SaveSession() error = %v", err) + } + + auth, err := GetAuthByDID(did.String()) + if err != nil { + t.Fatalf("GetAuthByDID() error = %v", err) + } + if auth == nil { + t.Fatal("GetAuthByDID() = nil, want auth") + } + if auth.RefreshJWT != session.RefreshToken { + t.Fatalf("RefreshJWT = %q, want %q", auth.RefreshJWT, session.RefreshToken) + } + if auth.DPoPHostNonce != session.DPoPHostNonce { + t.Fatalf("DPoPHostNonce = %q, want %q", auth.DPoPHostNonce, session.DPoPHostNonce) + } + + got, err := store.GetSession(context.Background(), did, session.SessionID) + if err != nil { + t.Fatalf("GetSession() error = %v", err) + } + if got.AccessToken != session.AccessToken { + t.Fatalf("AccessToken = %q, want %q", got.AccessToken, session.AccessToken) + } + + if err := store.DeleteSession(context.Background(), did, session.SessionID); err != nil { + t.Fatalf("DeleteSession() error = %v", err) + } + + deleted, err := store.GetSession(context.Background(), did, session.SessionID) + if err == nil || deleted != nil { + t.Fatalf("GetSession() after delete = (%v, %v), want error", deleted, err) + } +} diff --git a/frontend/src/App.svelte b/frontend/src/App.svelte index 45b0e11..cb4f4a4 100644 --- a/frontend/src/App.svelte +++ b/frontend/src/App.svelte @@ -119,13 +119,9 @@ } async function performSearch(query: string, source: string) { - if (!query.trim()) { - query = "*"; - } - isSearching = true; try { - const results = await Search(query, source); + const results = await Search(query.trim(), source); searchResults = sortResults(results); } catch (err) { console.error("Search failed:", err); diff --git a/index_service.go b/index_service.go index db7bf2a..1153a9f 100644 --- a/index_service.go +++ b/index_service.go @@ -153,31 +153,8 @@ func (s *IndexService) createClient() (*BlueskyClient, error) { return nil, fmt.Errorf("invalid DID: %w", err) } - redirectURI := "http://127.0.0.1/callback" - scopes := []string{"atproto", "transition:generic"} - config := oauth.NewLocalhostConfig(redirectURI, scopes) - store := oauth.NewMemStore() - - sessionData := oauth.ClientSessionData{ - AccountDID: did, - SessionID: auth.SessionID, - HostURL: auth.PDSURL, - AuthServerURL: auth.AuthServerURL, - AuthServerTokenEndpoint: auth.AuthServerTokenEndpoint, - AuthServerRevocationEndpoint: auth.AuthServerRevocationEndpoint, - AccessToken: auth.AccessJWT, - RefreshToken: auth.RefreshJWT, - Scopes: scopes, - DPoPAuthServerNonce: auth.DPoPAuthNonce, - DPoPHostNonce: auth.DPoPHostNonce, - DPoPPrivateKeyMultibase: auth.DPoPPrivateKey, - } - - if err := store.SaveSession(ctx, sessionData); err != nil { - return nil, fmt.Errorf("failed to save session: %w", err) - } - - app := oauth.NewClientApp(&config, store) + store := NewSQLiteOAuthStore() + app := newOAuthApp(store) session, err := app.ResumeSession(ctx, did, auth.SessionID) if err != nil { @@ -241,7 +218,7 @@ type BlueskyClient struct { } // fetchBookmarks writes bookmarks to the provided channel in batches -func (c *BlueskyClient) fetchBookmarks(maxPosts int, ch chan<- *PostResult, svc *IndexService) { +func (c *BlueskyClient) fetchBookmarks(maxPosts int, ch chan<- *PostResult, _ *IndexService) { ctx := context.Background() apiClient := c.session.APIClient() var cursor string @@ -291,7 +268,7 @@ func (c *BlueskyClient) fetchBookmarks(maxPosts int, ch chan<- *PostResult, svc } // fetchLikes writes likes to the provided channel in batches -func (c *BlueskyClient) fetchLikes(maxPosts int, ch chan<- *PostResult, svc *IndexService) { +func (c *BlueskyClient) fetchLikes(maxPosts int, ch chan<- *PostResult, _ *IndexService) { ctx := context.Background() apiClient := c.session.APIClient() var cursor string @@ -395,7 +372,7 @@ type postRecord struct { } // parsePostRecord extracts post data and facets from the LexiconTypeDecoder -func (c *BlueskyClient) parsePostRecord(decoder interface{}) (*postRecord, string, error) { +func (c *BlueskyClient) parsePostRecord(decoder any) (*postRecord, string, error) { if decoder == nil { return &postRecord{Text: "", CreatedAt: ""}, "", nil } diff --git a/oauth_store.go b/oauth_store.go new file mode 100644 index 0000000..58a2cb2 --- /dev/null +++ b/oauth_store.go @@ -0,0 +1,163 @@ +package main + +import ( + "context" + "fmt" + "net/http" + "sync" + "time" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" +) + +const oauthCallbackPort = 8787 + +var oauthScopes = []string{"atproto", "transition:generic"} + +func oauthCallbackURL() string { + return fmt.Sprintf("http://127.0.0.1:%d/callback", oauthCallbackPort) +} + +func oauthConfig() oauth.ClientConfig { + return oauth.NewLocalhostConfig(oauthCallbackURL(), append([]string(nil), oauthScopes...)) +} + +func newOAuthApp(store oauth.ClientAuthStore) *oauth.ClientApp { + config := oauthConfig() + return oauth.NewClientApp(&config, store) +} + +func authFromSessionData(sess *oauth.ClientSessionData, handle string) *Auth { + if handle == "" { + handle = sess.AccountDID.String() + } + + return &Auth{ + DID: sess.AccountDID.String(), + Handle: handle, + AccessJWT: sess.AccessToken, + RefreshJWT: sess.RefreshToken, + PDSURL: sess.HostURL, + SessionID: sess.SessionID, + AuthServerURL: sess.AuthServerURL, + AuthServerTokenEndpoint: sess.AuthServerTokenEndpoint, + AuthServerRevocationEndpoint: sess.AuthServerRevocationEndpoint, + DPoPAuthNonce: sess.DPoPAuthServerNonce, + DPoPHostNonce: sess.DPoPHostNonce, + DPoPPrivateKey: sess.DPoPPrivateKeyMultibase, + UpdatedAt: time.Now(), + } +} + +func sessionDataFromAuth(auth *Auth) (*oauth.ClientSessionData, error) { + did, err := syntax.ParseDID(auth.DID) + if err != nil { + return nil, fmt.Errorf("invalid DID in database: %w", err) + } + + return &oauth.ClientSessionData{ + AccountDID: did, + SessionID: auth.SessionID, + HostURL: auth.PDSURL, + AuthServerURL: auth.AuthServerURL, + AuthServerTokenEndpoint: auth.AuthServerTokenEndpoint, + AuthServerRevocationEndpoint: auth.AuthServerRevocationEndpoint, + Scopes: append([]string(nil), oauthScopes...), + AccessToken: auth.AccessJWT, + RefreshToken: auth.RefreshJWT, + DPoPAuthServerNonce: auth.DPoPAuthNonce, + DPoPHostNonce: auth.DPoPHostNonce, + DPoPPrivateKeyMultibase: auth.DPoPPrivateKey, + }, nil +} + +type SQLiteOAuthStore struct { + requests map[string]oauth.AuthRequestData + mu sync.Mutex +} + +func NewSQLiteOAuthStore() *SQLiteOAuthStore { + return &SQLiteOAuthStore{ + requests: make(map[string]oauth.AuthRequestData), + } +} + +func (s *SQLiteOAuthStore) GetSession(ctx context.Context, did syntax.DID, sessionID string) (*oauth.ClientSessionData, error) { + auth, err := GetAuthByDID(did.String()) + if err != nil { + return nil, err + } + if auth == nil || auth.SessionID != sessionID { + return nil, fmt.Errorf("session not found: %s", did) + } + + return sessionDataFromAuth(auth) +} + +func (s *SQLiteOAuthStore) SaveSession(ctx context.Context, sess oauth.ClientSessionData) error { + auth, err := GetAuthByDID(sess.AccountDID.String()) + if err != nil { + return err + } + + handle := "" + if auth != nil { + handle = auth.Handle + } + + return UpsertAuth(authFromSessionData(&sess, handle)) +} + +func (s *SQLiteOAuthStore) DeleteSession(ctx context.Context, did syntax.DID, sessionID string) error { + _, err := db.ExecContext(ctx, "DELETE FROM auth WHERE did = ? AND session_id = ?", did.String(), sessionID) + return err +} + +func (s *SQLiteOAuthStore) GetAuthRequestInfo(ctx context.Context, state string) (*oauth.AuthRequestData, error) { + s.mu.Lock() + defer s.mu.Unlock() + + info, ok := s.requests[state] + if !ok { + return nil, fmt.Errorf("request info not found: %s", state) + } + return &info, nil +} + +func (s *SQLiteOAuthStore) SaveAuthRequestInfo(ctx context.Context, info oauth.AuthRequestData) error { + s.mu.Lock() + defer s.mu.Unlock() + + if _, ok := s.requests[info.State]; ok { + return fmt.Errorf("auth request already saved for state %s", info.State) + } + + s.requests[info.State] = info + return nil +} + +func (s *SQLiteOAuthStore) DeleteAuthRequestInfo(ctx context.Context, state string) error { + s.mu.Lock() + defer s.mu.Unlock() + + delete(s.requests, state) + return nil +} + +func listenerAddress() string { + return fmt.Sprintf("127.0.0.1:%d", oauthCallbackPort) +} + +func closeCallbackServer(server *http.Server, listener httpCloser) { + if server != nil { + _ = server.Close() + } + if listener != nil { + _ = listener.Close() + } +} + +type httpCloser interface { + Close() error +}