diff --git a/internal/app/app.go b/internal/app/app.go index ebf9e35..d1c0085 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -1659,9 +1659,16 @@ func (a *App) health(w http.ResponseWriter, r *http.Request) { } func (a *App) pollLoop() { + changes := a.client.Changes() + a.poll() + ticker := time.NewTicker(a.interval) + defer ticker.Stop() for { + select { + case <-ticker.C: + case <-changes: + } a.poll() - time.Sleep(a.interval) } } diff --git a/internal/electrum/client.go b/internal/electrum/client.go index 2291fe3..31a604f 100644 --- a/internal/electrum/client.go +++ b/internal/electrum/client.go @@ -12,6 +12,7 @@ import ( "errors" "fmt" "net" + "sync" "sync/atomic" "time" @@ -21,6 +22,25 @@ import ( type Client struct { Address string nextID atomic.Uint64 + + mu sync.Mutex + session *session + + subscriptionsMu sync.Mutex + subscriptions map[string]*subscription + changes chan struct{} +} + +type subscription struct { + fetchMu sync.Mutex + mu sync.Mutex + + session *session + status string + generation uint64 + dirty bool + initialized bool + snapshot Snapshot } type HistoryItem struct { @@ -60,24 +80,111 @@ type Effect struct { } type response struct { - ID uint64 `json:"id"` - Result json.RawMessage `json:"result"` + ID *uint64 `json:"id,omitempty"` + Method string `json:"method,omitempty"` + Params []json.RawMessage `json:"params,omitempty"` + Result json.RawMessage `json:"result"` Error *struct { Code int `json:"code"` Message string `json:"message"` } `json:"error"` } +type rpcError struct { + code int + message string +} + +func (e rpcError) Error() string { + return fmt.Sprintf("electrum error %d: %s", e.code, e.message) +} + +type callResult struct { + result json.RawMessage + err error +} + +type session struct { + conn net.Conn + + writeMu sync.Mutex + pendingMu sync.Mutex + pending map[uint64]chan callResult + + done chan struct{} + closeOnce sync.Once + onNotify func(*session, response) +} + func (c *Client) Snapshot(ctx context.Context, scriptHash string) (Snapshot, error) { - var history []HistoryItem - if err := c.call(ctx, "blockchain.scripthash.get_history", []any{scriptHash}, &history); err != nil { - return Snapshot{}, err + sub := c.subscription(scriptHash) + sub.fetchMu.Lock() + defer sub.fetchMu.Unlock() + + var lastErr error + for attempt := 0; attempt < 2; attempt++ { + s, err := c.getSession(ctx) + if err != nil { + return Snapshot{}, err + } + if err := c.ensureSubscribed(ctx, s, scriptHash, sub); err != nil { + lastErr = err + if !retryable(err) { + return Snapshot{}, err + } + c.invalidate(s, err) + continue + } + + sub.mu.Lock() + if sub.initialized && !sub.dirty { + snapshot := cloneSnapshot(sub.snapshot) + sub.mu.Unlock() + return snapshot, nil + } + generation := sub.generation + sub.mu.Unlock() + + var history []HistoryItem + if err := s.call(ctx, c.nextID.Add(1), "blockchain.scripthash.get_history", []any{scriptHash}, &history); err != nil { + lastErr = err + if !retryable(err) { + return Snapshot{}, err + } + c.invalidate(s, err) + continue + } + var balance Balance + if err := s.call(ctx, c.nextID.Add(1), "blockchain.scripthash.get_balance", []any{scriptHash}, &balance); err != nil { + lastErr = err + if !retryable(err) { + return Snapshot{}, err + } + c.invalidate(s, err) + continue + } + + snapshot := Snapshot{Balance: balance, History: history} + sub.mu.Lock() + sub.snapshot = cloneSnapshot(snapshot) + sub.initialized = true + sub.dirty = sub.generation != generation + sub.mu.Unlock() + return snapshot, nil } - var balance Balance - if err := c.call(ctx, "blockchain.scripthash.get_balance", []any{scriptHash}, &balance); err != nil { - return Snapshot{}, err + return Snapshot{}, lastErr +} + +// Changes reports script-hash subscription updates. The channel is buffered +// and deliberately coalesces bursts so callers can refresh all dirty watches +// once instead of starting one scan per Electrum notification. +func (c *Client) Changes() <-chan struct{} { + c.mu.Lock() + defer c.mu.Unlock() + if c.changes == nil { + c.changes = make(chan struct{}, 1) } - return Snapshot{Balance: balance, History: history}, nil + return c.changes } func (c *Client) Ping(ctx context.Context) error { @@ -205,39 +312,242 @@ func (c *Client) transaction(ctx context.Context, txID string) (bitcoin.Transact } func (c *Client) call(ctx context.Context, method string, params []any, out any) error { + var lastErr error + for attempt := 0; attempt < 2; attempt++ { + s, err := c.getSession(ctx) + if err != nil { + return err + } + err = s.call(ctx, c.nextID.Add(1), method, params, out) + if err == nil || !retryable(err) { + return err + } + lastErr = err + c.invalidate(s, err) + } + return lastErr +} + +func (c *Client) getSession(ctx context.Context) (*session, error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.session != nil && !c.session.closed() { + return c.session, nil + } + dialer := net.Dialer{Timeout: 5 * time.Second} conn, err := dialer.DialContext(ctx, "tcp", c.Address) if err != nil { - return fmt.Errorf("connect to electrs: %w", err) + return nil, fmt.Errorf("connect to electrs: %w", err) } - defer conn.Close() - deadline := time.Now().Add(10 * time.Second) - if d, ok := ctx.Deadline(); ok && d.Before(deadline) { - deadline = d + s := &session{ + conn: conn, + pending: make(map[uint64]chan callResult), + done: make(chan struct{}), + onNotify: c.handleNotification, + } + c.session = s + go s.readLoop() + return s, nil +} + +func (c *Client) invalidate(s *session, err error) { + c.mu.Lock() + if c.session == s { + c.session = nil + } + c.mu.Unlock() + s.close(err) +} + +func (c *Client) subscription(scriptHash string) *subscription { + c.subscriptionsMu.Lock() + defer c.subscriptionsMu.Unlock() + if c.subscriptions == nil { + c.subscriptions = make(map[string]*subscription) + } + if c.subscriptions[scriptHash] == nil { + c.subscriptions[scriptHash] = &subscription{dirty: true} + } + return c.subscriptions[scriptHash] +} + +func (c *Client) ensureSubscribed(ctx context.Context, s *session, scriptHash string, sub *subscription) error { + sub.mu.Lock() + alreadySubscribed := sub.session == s + sub.mu.Unlock() + if alreadySubscribed { + return nil + } + + var status *string + if err := s.call(ctx, c.nextID.Add(1), "blockchain.scripthash.subscribe", []any{scriptHash}, &status); err != nil { + return err + } + statusValue := "" + if status != nil { + statusValue = *status + } + sub.mu.Lock() + if sub.session != s || sub.status != statusValue { + sub.generation++ + if sub.initialized && sub.status != statusValue { + sub.dirty = true + } + } + sub.session = s + sub.status = statusValue + sub.mu.Unlock() + return nil +} + +func (c *Client) handleNotification(s *session, message response) { + if message.Method != "blockchain.scripthash.subscribe" || len(message.Params) < 2 { + return + } + var scriptHash string + var status *string + if json.Unmarshal(message.Params[0], &scriptHash) != nil || json.Unmarshal(message.Params[1], &status) != nil || scriptHash == "" { + return + } + statusValue := "" + if status != nil { + statusValue = *status + } + sub := c.subscription(scriptHash) + sub.mu.Lock() + if sub.session != s || sub.status == statusValue { + sub.mu.Unlock() + return + } + sub.status = statusValue + sub.generation++ + sub.dirty = true + sub.mu.Unlock() + + c.mu.Lock() + changes := c.changes + c.mu.Unlock() + if changes != nil { + select { + case changes <- struct{}{}: + default: + } + } +} + +func cloneSnapshot(snapshot Snapshot) Snapshot { + clone := snapshot + clone.History = append([]HistoryItem(nil), snapshot.History...) + return clone +} + +func retryable(err error) bool { + var rpcErr rpcError + return err != nil && !errors.As(err, &rpcErr) && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) +} + +func (s *session) call(ctx context.Context, id uint64, method string, params []any, out any) error { + result := make(chan callResult, 1) + s.pendingMu.Lock() + if s.closed() { + s.pendingMu.Unlock() + return errors.New("electrum connection is closed") } - _ = conn.SetDeadline(deadline) + s.pending[id] = result + s.pendingMu.Unlock() - id := c.nextID.Add(1) req := map[string]any{"jsonrpc": "2.0", "id": id, "method": method, "params": params} - if err := json.NewEncoder(conn).Encode(req); err != nil { - return fmt.Errorf("write electrum request: %w", err) + s.writeMu.Lock() + deadline := time.Now().Add(10 * time.Second) + if d, ok := ctx.Deadline(); ok && d.Before(deadline) { + deadline = d } - line, err := bufio.NewReader(conn).ReadBytes('\n') + _ = s.conn.SetWriteDeadline(deadline) + err := json.NewEncoder(s.conn).Encode(req) + _ = s.conn.SetWriteDeadline(time.Time{}) + s.writeMu.Unlock() if err != nil { - return fmt.Errorf("read electrum response: %w", err) + s.removePending(id) + s.close(fmt.Errorf("write electrum request: %w", err)) + return fmt.Errorf("write electrum request: %w", err) } - var res response - if err := json.Unmarshal(line, &res); err != nil { - return fmt.Errorf("decode electrum response: %w", err) + + select { + case reply := <-result: + if reply.err != nil { + return reply.err + } + if err := json.Unmarshal(reply.result, out); err != nil { + return fmt.Errorf("decode electrum result: %w", err) + } + return nil + case <-ctx.Done(): + s.removePending(id) + return ctx.Err() + case <-s.done: + s.removePending(id) + return errors.New("electrum connection closed") } - if res.ID != id { - return errors.New("electrum returned an unexpected response id") +} + +func (s *session) readLoop() { + decoder := json.NewDecoder(bufio.NewReader(s.conn)) + for { + var message response + if err := decoder.Decode(&message); err != nil { + s.close(fmt.Errorf("read electrum response: %w", err)) + return + } + if message.ID == nil { + if s.onNotify != nil && message.Method != "" { + s.onNotify(s, message) + } + continue + } + s.pendingMu.Lock() + pending := s.pending[*message.ID] + delete(s.pending, *message.ID) + s.pendingMu.Unlock() + if pending == nil { + continue + } + if message.Error != nil { + pending <- callResult{err: rpcError{code: message.Error.Code, message: message.Error.Message}} + } else { + pending <- callResult{result: message.Result} + } } - if res.Error != nil { - return fmt.Errorf("electrum error %d: %s", res.Error.Code, res.Error.Message) +} + +func (s *session) removePending(id uint64) { + s.pendingMu.Lock() + delete(s.pending, id) + s.pendingMu.Unlock() +} + +func (s *session) close(err error) { + if err == nil { + err = errors.New("electrum connection closed") } - if err := json.Unmarshal(res.Result, out); err != nil { - return fmt.Errorf("decode electrum result: %w", err) + s.closeOnce.Do(func() { + _ = s.conn.Close() + close(s.done) + s.pendingMu.Lock() + pending := s.pending + s.pending = make(map[uint64]chan callResult) + s.pendingMu.Unlock() + for _, waiter := range pending { + waiter <- callResult{err: err} + } + }) +} + +func (s *session) closed() bool { + select { + case <-s.done: + return true + default: + return false } - return nil } diff --git a/internal/electrum/client_test.go b/internal/electrum/client_test.go index 849c94e..59c267a 100644 --- a/internal/electrum/client_test.go +++ b/internal/electrum/client_test.go @@ -4,8 +4,13 @@ package electrum import ( + "bufio" + "context" "encoding/binary" "encoding/hex" + "encoding/json" + "net" + "sync" "testing" "time" @@ -69,3 +74,264 @@ func TestAddAmountSumsRepeatedOutputs(t *testing.T) { t.Fatalf("repeated output amounts were not summed: %#v", amounts) } } + +func TestSnapshotSubscribesCachesAndRefreshesOnNotification(t *testing.T) { + server := newTestElectrumServer(t) + client := &Client{Address: server.address()} + changes := client.Changes() + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + first, err := client.Snapshot(ctx, "script-hash") + if err != nil { + t.Fatal(err) + } + if first.Balance.Confirmed != 42 { + t.Fatalf("initial balance = %d, want 42", first.Balance.Confirmed) + } + + second, err := client.Snapshot(ctx, "script-hash") + if err != nil { + t.Fatal(err) + } + if second.Balance.Confirmed != 42 { + t.Fatalf("cached balance = %d, want 42", second.Balance.Confirmed) + } + if got := server.methodCount("blockchain.scripthash.subscribe"); got != 1 { + t.Fatalf("subscribe calls = %d, want 1", got) + } + if got := server.methodCount("blockchain.scripthash.get_history"); got != 1 { + t.Fatalf("history calls before a change = %d, want 1", got) + } + if got := server.methodCount("blockchain.scripthash.get_balance"); got != 1 { + t.Fatalf("balance calls before a change = %d, want 1", got) + } + + server.setState("changed", 84) + server.notify(t, "script-hash", "changed") + select { + case <-changes: + case <-ctx.Done(): + t.Fatal("subscription change was not reported") + } + + third, err := client.Snapshot(ctx, "script-hash") + if err != nil { + t.Fatal(err) + } + if third.Balance.Confirmed != 84 { + t.Fatalf("refreshed balance = %d, want 84", third.Balance.Confirmed) + } + if got := server.methodCount("blockchain.scripthash.get_history"); got != 2 { + t.Fatalf("history calls after a change = %d, want 2", got) + } + if got := server.methodCount("blockchain.scripthash.get_balance"); got != 2 { + t.Fatalf("balance calls after a change = %d, want 2", got) + } + if violations := server.unsubscribedQueries(); len(violations) != 0 { + t.Fatalf("queries were sent before subscribing on their connection: %v", violations) + } +} + +func TestSnapshotReconnectsAndResubscribes(t *testing.T) { + server := newTestElectrumServer(t) + client := &Client{Address: server.address()} + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + if _, err := client.Snapshot(ctx, "script-hash"); err != nil { + t.Fatal(err) + } + server.setState("after-reconnect", 99) + server.disconnect(t) + + deadline := time.Now().Add(time.Second) + for { + client.mu.Lock() + closed := client.session == nil || client.session.closed() + client.mu.Unlock() + if closed { + break + } + if time.Now().After(deadline) { + t.Fatal("client did not observe the disconnected session") + } + time.Sleep(time.Millisecond) + } + + snapshot, err := client.Snapshot(ctx, "script-hash") + if err != nil { + t.Fatal(err) + } + if snapshot.Balance.Confirmed != 99 { + t.Fatalf("balance after reconnect = %d, want 99", snapshot.Balance.Confirmed) + } + if got := server.methodCount("blockchain.scripthash.subscribe"); got != 2 { + t.Fatalf("subscribe calls after reconnect = %d, want 2", got) + } + if violations := server.unsubscribedQueries(); len(violations) != 0 { + t.Fatalf("queries were sent before resubscribing: %v", violations) + } +} + +type testElectrumServer struct { + listener net.Listener + + mu sync.Mutex + connections []net.Conn + current net.Conn + methods map[string]int + violations []string + status string + balance int64 + + writeMu sync.Mutex +} + +func newTestElectrumServer(t *testing.T) *testElectrumServer { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + server := &testElectrumServer{ + listener: listener, + methods: make(map[string]int), + status: "initial", + balance: 42, + } + go server.acceptLoop() + t.Cleanup(func() { + _ = listener.Close() + server.mu.Lock() + connections := append([]net.Conn(nil), server.connections...) + server.mu.Unlock() + for _, conn := range connections { + _ = conn.Close() + } + }) + return server +} + +func (s *testElectrumServer) address() string { + return s.listener.Addr().String() +} + +func (s *testElectrumServer) acceptLoop() { + for { + conn, err := s.listener.Accept() + if err != nil { + return + } + s.mu.Lock() + s.connections = append(s.connections, conn) + s.current = conn + s.mu.Unlock() + go s.serve(conn) + } +} + +func (s *testElectrumServer) serve(conn net.Conn) { + subscribed := make(map[string]bool) + decoder := json.NewDecoder(bufio.NewReader(conn)) + for { + var request struct { + ID uint64 `json:"id"` + Method string `json:"method"` + Params []json.RawMessage `json:"params"` + } + if decoder.Decode(&request) != nil { + return + } + var scriptHash string + if len(request.Params) > 0 { + _ = json.Unmarshal(request.Params[0], &scriptHash) + } + + s.mu.Lock() + s.methods[request.Method]++ + status, balance := s.status, s.balance + if (request.Method == "blockchain.scripthash.get_history" || request.Method == "blockchain.scripthash.get_balance") && !subscribed[scriptHash] { + s.violations = append(s.violations, request.Method) + } + s.mu.Unlock() + + var result any + switch request.Method { + case "blockchain.scripthash.subscribe": + subscribed[scriptHash] = true + result = status + case "blockchain.scripthash.get_history": + result = []HistoryItem{{TxHash: "tx", Height: 1}} + case "blockchain.scripthash.get_balance": + result = Balance{Confirmed: balance} + default: + result = nil + } + s.writeMu.Lock() + err := json.NewEncoder(conn).Encode(map[string]any{"jsonrpc": "2.0", "id": request.ID, "result": result}) + s.writeMu.Unlock() + if err != nil { + return + } + } +} + +func (s *testElectrumServer) setState(status string, balance int64) { + s.mu.Lock() + s.status = status + s.balance = balance + s.mu.Unlock() +} + +func (s *testElectrumServer) notify(t *testing.T, scriptHash, status string) { + t.Helper() + deadline := time.Now().Add(time.Second) + for { + s.mu.Lock() + conn := s.current + s.mu.Unlock() + if conn != nil { + s.writeMu.Lock() + err := json.NewEncoder(conn).Encode(map[string]any{ + "jsonrpc": "2.0", + "method": "blockchain.scripthash.subscribe", + "params": []any{scriptHash, status}, + }) + s.writeMu.Unlock() + if err != nil { + t.Fatal(err) + } + return + } + if time.Now().After(deadline) { + t.Fatal("test server did not receive a connection") + } + time.Sleep(time.Millisecond) + } +} + +func (s *testElectrumServer) disconnect(t *testing.T) { + t.Helper() + s.mu.Lock() + conn := s.current + s.mu.Unlock() + if conn == nil { + t.Fatal("test server has no connection to close") + } + if err := conn.Close(); err != nil { + t.Fatal(err) + } +} + +func (s *testElectrumServer) methodCount(method string) int { + s.mu.Lock() + defer s.mu.Unlock() + return s.methods[method] +} + +func (s *testElectrumServer) unsubscribedQueries() []string { + s.mu.Lock() + defer s.mu.Unlock() + return append([]string(nil), s.violations...) +}