diff --git a/backend/radiance.go b/backend/radiance.go index 39fecf6a..04596ccc 100644 --- a/backend/radiance.go +++ b/backend/radiance.go @@ -21,6 +21,7 @@ import ( "go.opentelemetry.io/otel/trace" C "github.com/getlantern/common" + wire "github.com/getlantern/common/usermessage" "github.com/getlantern/publicip" "github.com/getlantern/radiance/account" @@ -38,6 +39,7 @@ import ( "github.com/getlantern/radiance/telemetry" "github.com/getlantern/radiance/traces" "github.com/getlantern/radiance/unbounded" + clientmessage "github.com/getlantern/radiance/usermessage" "github.com/getlantern/radiance/vpn" lbA "github.com/getlantern/lantern-box/adapter" @@ -57,6 +59,7 @@ type LocalBackend struct { confHandler *config.ConfigHandler issueReporter *issue.IssueReporter accountClient *account.Client + userMessages *clientmessage.Service srvManager *servers.Manager vpnClient *vpn.VPNClient @@ -227,6 +230,27 @@ func NewLocalBackend(ctx context.Context, opts Options) (*LocalBackend, error) { } r.sessionHistory = vpn.NewSessionHistory(slog.Default().With("service", "session_history"), r.sessionInfo()) r.shutdownFuncs = append(r.shutdownFuncs, func() error { r.sessionHistory.Close(); return nil }) + userMessages, err := clientmessage.New(clientmessage.Options{ + DataDir: dataDir, + Fetcher: clientmessage.NewHTTPFetcher( + kindling.HTTPClient(), + clientmessage.Endpoint(common.GetBaseURL()), + ), + ContextProvider: func() clientmessage.ClientContext { + return clientmessage.ClientContext{ + UserID: settings.GetString(settings.UserIDKey), + ProToken: settings.GetString(settings.TokenKey), + Locale: clientmessage.NormalizeLocale(settings.GetString(settings.LocaleKey)), + Platform: clientmessage.NormalizePlatform(common.Platform), + AppVersion: common.GetVersion(), + } + }, + }) + if err != nil { + slog.Error("Loading user-message state", "error", err) + } else { + r.userMessages = userMessages + } r.clearSelectedIfMissing() return r, nil } @@ -234,6 +258,12 @@ func NewLocalBackend(ctx context.Context, opts Options) (*LocalBackend, error) { func (r *LocalBackend) Start() { // eagerly start kindling so it's ready by the time we need to make network requests kindling.Init() + if r.userMessages != nil { + events.SubscribeContext(r.ctx, func(account.UserChangeEvent) { + r.userMessages.Refresh() + }) + r.userMessages.Start(r.ctx) + } go func() { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) result, err := publicip.Detect(ctx, &publicip.Config{ @@ -601,6 +631,9 @@ func (r *LocalBackend) PatchSettings(updates settings.Settings) error { if err := settings.Patch(diff); err != nil { return fmt.Errorf("failed to update settings: %w", err) } + if _, ok := diff[settings.LocaleKey]; ok { + r.refreshUserMessages() + } // telemetry settings if _, ok := diff[settings.TelemetryKey]; ok { if settings.GetBool(settings.TelemetryKey) { @@ -1548,19 +1581,35 @@ func (r *LocalBackend) RemoveSplitTunnelItems(items vpn.SplitTunnelFilter) error ///////////// func (r *LocalBackend) NewUser(ctx context.Context) (*account.UserData, error) { - return r.accountClient.NewUser(ctx) + userData, err := r.accountClient.NewUser(ctx) + if err == nil { + r.refreshUserMessages() + } + return userData, err } func (r *LocalBackend) Login(ctx context.Context, email, password string) (*account.UserData, error) { - return r.accountClient.Login(ctx, email, password) + userData, err := r.accountClient.Login(ctx, email, password) + if err == nil { + r.refreshUserMessages() + } + return userData, err } func (r *LocalBackend) Logout(ctx context.Context, email string) (*account.UserData, error) { - return r.accountClient.Logout(ctx, email) + userData, err := r.accountClient.Logout(ctx, email) + if err == nil { + r.refreshUserMessages() + } + return userData, err } func (r *LocalBackend) FetchUserData(ctx context.Context) (*account.UserData, error) { - return r.accountClient.FetchUserData(ctx) + userData, err := r.accountClient.FetchUserData(ctx) + if err == nil { + r.refreshUserMessages() + } + return userData, err } func (r *LocalBackend) VerifyPassword(ctx context.Context, email, password string) error { @@ -1585,7 +1634,11 @@ func (r *LocalBackend) CompleteRecoveryByEmail(ctx context.Context, email, newPa } func (r *LocalBackend) DeleteAccount(ctx context.Context, email, password string) (*account.UserData, error) { - return r.accountClient.DeleteAccount(ctx, email, password) + userData, err := r.accountClient.DeleteAccount(ctx, email, password) + if err == nil { + r.refreshUserMessages() + } + return userData, err } func (r *LocalBackend) SignUp(ctx context.Context, email, password string) ([]byte, *account.SignupResponse, error) { @@ -1652,7 +1705,11 @@ func (r *LocalBackend) RemoveDevice(ctx context.Context, deviceID string) (*acco } func (r *LocalBackend) OAuthLoginCallback(ctx context.Context, oAuthToken string) (*account.UserData, error) { - return r.accountClient.OAuthLoginCallback(ctx, oAuthToken) + userData, err := r.accountClient.OAuthLoginCallback(ctx, oAuthToken) + if err == nil { + r.refreshUserMessages() + } + return userData, err } func (r *LocalBackend) OAuthLoginURL(ctx context.Context, provider string) (string, error) { @@ -1671,6 +1728,40 @@ func (r *LocalBackend) UserData() (*account.UserData, error) { return &userData, nil } +// CurrentUserMessage returns the current account's pending message. +func (r *LocalBackend) CurrentUserMessage() (*wire.ResolvedUserMessage, error) { + if r.userMessages == nil { + return nil, nil + } + return r.userMessages.Current() +} + +// RefreshUserMessages schedules an immediate eligibility refresh. +func (r *LocalBackend) RefreshUserMessages() { + r.refreshUserMessages() +} + +func (r *LocalBackend) refreshUserMessages() { + if r.userMessages != nil { + r.userMessages.Refresh() + } +} + +// AcknowledgeUserMessage records that the UI displayed displayID. +func (r *LocalBackend) AcknowledgeUserMessage(displayID string) error { + if r.userMessages == nil { + return clientmessage.ErrMessageNotPending + } + return r.userMessages.Acknowledge(displayID) +} + +// SetUserMessageActivity adjusts polling for app and network lifecycle changes. +func (r *LocalBackend) SetUserMessageActivity(active, online bool) { + if r.userMessages != nil { + r.userMessages.SetActivity(active, online) + } +} + /////////////////// // Subscriptions // /////////////////// diff --git a/go.mod b/go.mod index f0c18666..8aa9eadb 100644 --- a/go.mod +++ b/go.mod @@ -40,7 +40,7 @@ require ( github.com/alitto/pond v1.9.2 github.com/getlantern/amp v0.0.0-20260606002220-a8629924577c github.com/getlantern/broflake v0.0.0-20260810172605-bef5e5234952 - github.com/getlantern/common v1.2.1-0.20260708083946-cc657b08792c + github.com/getlantern/common v1.2.1-0.20260818065623-10c2257aa54f github.com/getlantern/dnstt v0.0.0-20260603191204-3b860502c0ac github.com/getlantern/domainfront v0.0.0-20260722204513-8c1f8acfa715 github.com/getlantern/keepcurrent v0.0.0-20260616120552-f204338b01a3 @@ -67,6 +67,7 @@ require ( go.opentelemetry.io/otel/sdk v1.43.0 go.opentelemetry.io/otel/sdk/metric v1.43.0 golang.org/x/term v0.41.0 + golang.org/x/text v0.35.0 golang.zx2c4.com/wireguard/wgctrl v0.0.0-20241231184526-a9ab2273dd10 google.golang.org/protobuf v1.36.11 gopkg.in/natefinch/lumberjack.v2 v2.2.1 diff --git a/go.sum b/go.sum index 2ab51cdd..8d71e636 100644 --- a/go.sum +++ b/go.sum @@ -238,8 +238,8 @@ github.com/getlantern/amp v0.0.0-20260606002220-a8629924577c h1:ZzxuhIWO295y4eE6 github.com/getlantern/amp v0.0.0-20260606002220-a8629924577c/go.mod h1:b5teAOFT+vpBqc2CHoz74QTEM15iOv1PLdQogwXcfrM= github.com/getlantern/broflake v0.0.0-20260810172605-bef5e5234952 h1:nD8iJ4IpaTq/09PhoOzmaGTwatOxsR41neSmbCY6RS8= github.com/getlantern/broflake v0.0.0-20260810172605-bef5e5234952/go.mod h1:1+1kCIg9Zj+2CgN+vl868AlAZVUItuN4wLhFur5QakA= -github.com/getlantern/common v1.2.1-0.20260708083946-cc657b08792c h1:Hpxu12ORnAcyYuIqV2yAqrmDAnTaY1Cd3yL7TvFB6ME= -github.com/getlantern/common v1.2.1-0.20260708083946-cc657b08792c/go.mod h1:eSSuV4bMPgQJnczBw+KWWqWNo1itzmVxC++qUBPRTt0= +github.com/getlantern/common v1.2.1-0.20260818065623-10c2257aa54f h1:8iZqf4mUkFXKmUnDj2Iej0WySAzLuQmHOX736b4pCYU= +github.com/getlantern/common v1.2.1-0.20260818065623-10c2257aa54f/go.mod h1:eSSuV4bMPgQJnczBw+KWWqWNo1itzmVxC++qUBPRTt0= github.com/getlantern/context v0.0.0-20220418194847-3d5e7a086201 h1:oEZYEpZo28Wdx+5FZo4aU7JFXu0WG/4wJWese5reQSA= github.com/getlantern/context v0.0.0-20220418194847-3d5e7a086201/go.mod h1:Y9WZUHEb+mpra02CbQ/QczLUe6f0Dezxaw5DCJlJQGo= github.com/getlantern/dnstt v0.0.0-20260603191204-3b860502c0ac h1:TMvkNgLVyIYAfu1dYrOuRPYMVY+cHEo1C0CYWhtSw2A= diff --git a/ipc/client.go b/ipc/client.go index 3cfa0480..0dfb2ba6 100644 --- a/ipc/client.go +++ b/ipc/client.go @@ -16,6 +16,7 @@ import ( "syscall" "time" + wire "github.com/getlantern/common/usermessage" box "github.com/getlantern/lantern-box" "github.com/getlantern/radiance/account" @@ -205,6 +206,35 @@ func (c *Client) UpdateConfig(ctx context.Context) error { return err } +// CurrentUserMessage returns the pending message for the current account. +func (c *Client) CurrentUserMessage(ctx context.Context) (*wire.ResolvedUserMessage, error) { + var response CurrentUserMessageResponse + if err := c.doJSON(ctx, http.MethodGet, userMessageEndpoint, nil, &response); err != nil { + return nil, err + } + return response.Message, nil +} + +// RefreshUserMessages schedules an immediate server eligibility refresh. +func (c *Client) RefreshUserMessages(ctx context.Context) error { + _, err := c.do(ctx, http.MethodPost, userMessageRefreshEndpoint, nil) + return err +} + +// AcknowledgeUserMessage records that the UI displayed displayID. +func (c *Client) AcknowledgeUserMessage(ctx context.Context, displayID string) error { + _, err := c.do(ctx, http.MethodPost, userMessageAcknowledgeEndpoint, + UserMessageAcknowledgeRequest{DisplayID: displayID}) + return err +} + +// SetUserMessageActivity updates the app and connectivity lifecycle used by polling. +func (c *Client) SetUserMessageActivity(ctx context.Context, active, online bool) error { + _, err := c.do(ctx, http.MethodPatch, userMessageActivityEndpoint, + UserMessageActivityRequest{Active: active, Online: online}) + return err +} + /////////////////////// // Server management // /////////////////////// diff --git a/ipc/server.go b/ipc/server.go index e42908a9..3f176426 100644 --- a/ipc/server.go +++ b/ipc/server.go @@ -24,6 +24,7 @@ import ( rlog "github.com/getlantern/radiance/log" "github.com/getlantern/radiance/peer" "github.com/getlantern/radiance/unbounded" + clientmessage "github.com/getlantern/radiance/usermessage" "github.com/getlantern/radiance/vpn" sjson "github.com/sagernet/sing/common/json" @@ -54,6 +55,11 @@ const ( configEventsEndpoint = "/config/events" configUpdateEndpoint = "/config/update" + userMessageEndpoint = "/user-messages" + userMessageRefreshEndpoint = "/user-messages/refresh" + userMessageAcknowledgeEndpoint = "/user-messages/acknowledge" + userMessageActivityEndpoint = "/user-messages/activity" + // Server management endpoints serversEndpoint = "/servers" serversAddEndpoint = "/servers/add" @@ -229,6 +235,11 @@ func newLocalAPI(b *backend.LocalBackend, withAuth bool) *localapi { mux.HandleFunc("GET "+configEventsEndpoint, s.configEventsHandler) mux.HandleFunc("POST "+configUpdateEndpoint, traced(s.configUpdateHandler)) + mux.HandleFunc("GET "+userMessageEndpoint, traced(s.userMessageHandler)) + mux.HandleFunc("POST "+userMessageRefreshEndpoint, traced(s.userMessageRefreshHandler)) + mux.HandleFunc("POST "+userMessageAcknowledgeEndpoint, traced(s.userMessageAcknowledgeHandler)) + mux.HandleFunc("PATCH "+userMessageActivityEndpoint, traced(s.userMessageActivityHandler)) + // Server management mux.HandleFunc("GET "+serversEndpoint, traced(s.serversHandler)) mux.HandleFunc("POST "+serversAddEndpoint, traced(s.serversAddHandler)) @@ -896,6 +907,47 @@ func (s *localapi) settingsHandler(w http.ResponseWriter, r *http.Request) { } } +func (s *localapi) userMessageHandler(w http.ResponseWriter, r *http.Request) { + message, err := s.backend(r.Context()).CurrentUserMessage() + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + return + } + writeJSON(w, http.StatusOK, CurrentUserMessageResponse{Message: message}) +} + +func (s *localapi) userMessageRefreshHandler(w http.ResponseWriter, r *http.Request) { + s.backend(r.Context()).RefreshUserMessages() + w.WriteHeader(http.StatusNoContent) +} + +func (s *localapi) userMessageAcknowledgeHandler(w http.ResponseWriter, r *http.Request) { + var request UserMessageAcknowledgeRequest + if err := decodeJSON(r, &request); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + if err := s.backend(r.Context()).AcknowledgeUserMessage(request.DisplayID); err != nil { + status := http.StatusInternalServerError + if errors.Is(err, clientmessage.ErrMessageNotPending) { + status = http.StatusConflict + } + http.Error(w, err.Error(), status) + return + } + w.WriteHeader(http.StatusNoContent) +} + +func (s *localapi) userMessageActivityHandler(w http.ResponseWriter, r *http.Request) { + var request UserMessageActivityRequest + if err := decodeJSON(r, &request); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + s.backend(r.Context()).SetUserMessageActivity(request.Active, request.Online) + w.WriteHeader(http.StatusNoContent) +} + func (s *localapi) envHandler(w http.ResponseWriter, r *http.Request) { switch r.Method { case http.MethodPatch: diff --git a/ipc/types.go b/ipc/types.go index 88f215ad..198fe4d8 100644 --- a/ipc/types.go +++ b/ipc/types.go @@ -2,6 +2,7 @@ package ipc import ( "github.com/getlantern/common" + wire "github.com/getlantern/common/usermessage" "github.com/getlantern/radiance/account" "github.com/getlantern/radiance/issue" @@ -116,6 +117,17 @@ type IssueReportRequest struct { Attachments []*issue.Attachment `json:"attachments"` } +// UserMessageAcknowledgeRequest identifies the pending message displayed by the UI. +type UserMessageAcknowledgeRequest struct { + DisplayID string `json:"displayID"` +} + +// UserMessageActivityRequest reports whether foreground polling should run. +type UserMessageActivityRequest struct { + Active bool `json:"active"` + Online bool `json:"online"` +} + // Shared response types used by both client and server. type SelectedServerResponse struct { @@ -148,6 +160,11 @@ type PlansResponse struct { Plans string `json:"plans"` } +// CurrentUserMessageResponse contains the pending message, if one exists. +type CurrentUserMessageResponse struct { + Message *wire.ResolvedUserMessage `json:"message,omitempty"` +} + type ResultResponse struct { Result string `json:"result"` } diff --git a/ipc/usermessage_test.go b/ipc/usermessage_test.go new file mode 100644 index 00000000..8759062a --- /dev/null +++ b/ipc/usermessage_test.go @@ -0,0 +1,52 @@ +package ipc + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/getlantern/radiance/backend" +) + +func TestUserMessageRoutes(t *testing.T) { + api := newLocalAPI(&backend.LocalBackend{}, false) + + response := serveUserMessageRequest(t, api, http.MethodGet, userMessageEndpoint, nil) + require.Equal(t, http.StatusOK, response.Code) + var current CurrentUserMessageResponse + require.NoError(t, json.Unmarshal(response.Body.Bytes(), ¤t)) + require.Nil(t, current.Message) + + response = serveUserMessageRequest(t, api, http.MethodPost, userMessageRefreshEndpoint, nil) + require.Equal(t, http.StatusNoContent, response.Code) + + response = serveUserMessageRequest(t, api, http.MethodPatch, userMessageActivityEndpoint, + UserMessageActivityRequest{Active: true, Online: true}) + require.Equal(t, http.StatusNoContent, response.Code) + + response = serveUserMessageRequest(t, api, http.MethodPost, userMessageAcknowledgeEndpoint, + UserMessageAcknowledgeRequest{DisplayID: "not-pending"}) + require.Equal(t, http.StatusConflict, response.Code) +} + +func serveUserMessageRequest( + t *testing.T, + api http.Handler, + method string, + endpoint string, + body any, +) *httptest.ResponseRecorder { + t.Helper() + var encoded bytes.Buffer + if body != nil { + require.NoError(t, json.NewEncoder(&encoded).Encode(body)) + } + request := httptest.NewRequest(method, endpoint, &encoded) + response := httptest.NewRecorder() + api.ServeHTTP(response, request) + return response +} diff --git a/usermessage/context.go b/usermessage/context.go new file mode 100644 index 00000000..6b34c27a --- /dev/null +++ b/usermessage/context.go @@ -0,0 +1,25 @@ +package usermessage + +import ( + "strings" + + "golang.org/x/text/language" +) + +// NormalizeLocale returns a canonical BCP 47 tag with a safe fallback. +func NormalizeLocale(locale string) string { + tag, err := language.Parse(strings.TrimSpace(locale)) + if err != nil || tag == language.Und { + return "en-US" + } + return tag.String() +} + +// NormalizePlatform maps runtime platform names to the public targeting vocabulary. +func NormalizePlatform(platform string) string { + platform = strings.ToLower(strings.TrimSpace(platform)) + if platform == "darwin" { + return "macos" + } + return platform +} diff --git a/usermessage/http.go b/usermessage/http.go new file mode 100644 index 00000000..a3d5b47d --- /dev/null +++ b/usermessage/http.go @@ -0,0 +1,138 @@ +// Package usermessage retrieves and retains presentation-ready user messages. +package usermessage + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strconv" + "strings" + + wire "github.com/getlantern/common/usermessage" + "github.com/getlantern/kindling" + + "github.com/getlantern/radiance/common" +) + +const maxResponseBytes = 32 * 1024 + +var errCredentialsUnavailable = errors.New("user-message credentials are unavailable") + +// ClientContext is the authenticated and presentation context for one fetch. +type ClientContext struct { + UserID string + ProToken string + Locale string + Platform string + AppVersion string +} + +func (c ClientContext) valid() bool { + userID, err := strconv.ParseInt(c.UserID, 10, 64) + return err == nil && userID > 0 && strconv.FormatInt(userID, 10) == c.UserID && + c.ProToken != "" && len(c.ProToken) <= 4096 && + c.Locale != "" && c.Platform != "" && c.AppVersion != "" +} + +type httpStatusError struct { + statusCode int +} + +func (err *httpStatusError) Error() string { + return fmt.Sprintf("unexpected status %d", err.statusCode) +} + +// Fetcher resolves at most one message for a client context. +type Fetcher interface { + Fetch(context.Context, ClientContext, []string) (wire.UserMessageResponse, error) +} + +// HTTPFetcher implements Fetcher using Lantern Cloud's public endpoint. +type HTTPFetcher struct { + client *http.Client + endpoint string +} + +// NewHTTPFetcher creates a fetcher for endpoint. +func NewHTTPFetcher(client *http.Client, endpoint string) *HTTPFetcher { + return &HTTPFetcher{client: client, endpoint: endpoint} +} + +// Fetch requests one resolved message. Unsupported or otherwise unsafe messages +// are discarded while a valid polling recommendation is retained. +func (f *HTTPFetcher) Fetch( + ctx context.Context, + clientContext ClientContext, + seenDisplayIDs []string, +) (wire.UserMessageResponse, error) { + if !clientContext.valid() { + return wire.UserMessageResponse{}, errCredentialsUnavailable + } + request := wire.UserMessageRequest{ + Locale: clientContext.Locale, + Platform: clientContext.Platform, + AppVersion: clientContext.AppVersion, + Capability: wire.CapabilityUserMessagesV1, + SeenDisplayIDs: seenDisplayIDs, + } + if err := request.Validate(); err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("validate user-message request: %w", err) + } + body, err := json.Marshal(request) + if err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("marshal user-message request: %w", err) + } + req, err := common.NewRequestWithHeaders(ctx, http.MethodPost, f.endpoint, bytes.NewReader(body)) + if err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("create user-message request: %w", err) + } + req.Header.Set(common.ContentTypeHeader, "application/json") + req.Header.Set(common.AcceptHeader, "application/json") + req.Header.Set(common.UserIDHeader, clientContext.UserID) + req.Header.Set(common.ProTokenHeader, clientContext.ProToken) + req.Header.Set(common.PlatformHeader, clientContext.Platform) + req.Header.Set(common.AppVersionHeader, clientContext.AppVersion) + req.Header.Set(common.VersionHeader, clientContext.AppVersion) + req.Header.Set(common.AppNameHeader, common.Name) + req.Header.Set(kindling.IdempotentHeader, "1") + + resp, err := f.client.Do(req) + if err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("fetch user message: %w", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + return wire.UserMessageResponse{}, fmt.Errorf( + "fetch user message: %w", &httpStatusError{statusCode: resp.StatusCode}, + ) + } + data, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseBytes+1)) + if err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("read user-message response: %w", err) + } + if len(data) > maxResponseBytes { + return wire.UserMessageResponse{}, errors.New("user-message response exceeds size limit") + } + var response wire.UserMessageResponse + if err := json.Unmarshal(data, &response); err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("decode user-message response: %w", err) + } + pollOnly := response + pollOnly.Message = nil + if err := pollOnly.Validate(); err != nil { + return wire.UserMessageResponse{}, fmt.Errorf("validate user-message response: %w", err) + } + if response.Message != nil && response.Message.Validate() != nil { + response.Message = nil + } + return response, nil +} + +// Endpoint returns the public user-message endpoint under baseURL. +func Endpoint(baseURL string) string { + return strings.TrimRight(baseURL, "/") + "/user-messages" +} diff --git a/usermessage/http_test.go b/usermessage/http_test.go new file mode 100644 index 00000000..675c2b10 --- /dev/null +++ b/usermessage/http_test.go @@ -0,0 +1,135 @@ +package usermessage + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/stretchr/testify/require" + + wire "github.com/getlantern/common/usermessage" + "github.com/getlantern/kindling" + + "github.com/getlantern/radiance/common" +) + +func TestHTTPFetcherContractAndCredentials(t *testing.T) { + expiresAt := time.Now().Add(time.Hour).UTC().Truncate(time.Second) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + require.Equal(t, http.MethodPost, r.Method) + require.Equal(t, "/v1/user-messages", r.URL.Path) + require.Equal(t, "12345", r.Header.Get(common.UserIDHeader)) + require.Equal(t, "secret-token", r.Header.Get(common.ProTokenHeader)) + require.Equal(t, "macos", r.Header.Get(common.PlatformHeader)) + require.Equal(t, "9.2.0", r.Header.Get(common.AppVersionHeader)) + require.Equal(t, "1", r.Header.Get(kindling.IdempotentHeader)) + + var request wire.UserMessageRequest + require.NoError(t, json.NewDecoder(r.Body).Decode(&request)) + require.Equal(t, wire.CapabilityUserMessagesV1, request.Capability) + require.Equal(t, "fa-IR", request.Locale) + require.Equal(t, []string{"seen-1"}, request.SeenDisplayIDs) + + writeWireResponse(t, w, wire.UserMessageResponse{ + PollIntervalSeconds: wire.MaxPollIntervalSeconds, + Message: testMessage("display-1", expiresAt), + }) + })) + defer server.Close() + + fetcher := NewHTTPFetcher(server.Client(), server.URL+"/v1/user-messages") + response, err := fetcher.Fetch(context.Background(), testClientContext(), []string{"seen-1"}) + require.NoError(t, err) + require.Equal(t, "display-1", response.Message.DisplayID) +} + +func TestHTTPFetcherSafelyIgnoresUnsupportedMessages(t *testing.T) { + tests := map[string]func(*wire.ResolvedUserMessage){ + "surface": func(message *wire.ResolvedUserMessage) { + message.Surface = "future_surface" + }, + "action": func(message *wire.ResolvedUserMessage) { + message.ButtonLabel = "Act" + message.Action = &wire.Action{Type: "future_action"} + }, + } + for name, mutate := range tests { + t.Run(name, func(t *testing.T) { + message := testMessage("display-1", time.Now().Add(time.Hour)) + mutate(message) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + writeWireResponse(t, w, wire.UserMessageResponse{ + PollIntervalSeconds: 60, + Message: message, + }) + })) + defer server.Close() + + response, err := NewHTTPFetcher(server.Client(), server.URL).Fetch( + context.Background(), testClientContext(), nil, + ) + require.NoError(t, err) + require.Nil(t, response.Message) + require.Equal(t, 60, response.PollIntervalSeconds) + }) + } +} + +func TestHTTPFetcherRejectsNonCanonicalUserID(t *testing.T) { + for _, userID := range []string{"not-a-number", "00123", "+123", "9223372036854775808"} { + t.Run(userID, func(t *testing.T) { + clientContext := testClientContext() + clientContext.UserID = userID + _, err := NewHTTPFetcher(http.DefaultClient, "https://example.com").Fetch( + context.Background(), clientContext, nil, + ) + require.ErrorIs(t, err, errCredentialsUnavailable) + }) + } +} + +func TestHTTPFetcherReturnsStructuredStatusError(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + })) + defer server.Close() + + _, err := NewHTTPFetcher(server.Client(), server.URL).Fetch( + context.Background(), testClientContext(), nil, + ) + var statusErr *httpStatusError + require.ErrorAs(t, err, &statusErr) + require.Equal(t, http.StatusUnauthorized, statusErr.statusCode) +} + +func writeWireResponse(t *testing.T, w http.ResponseWriter, response wire.UserMessageResponse) { + t.Helper() + w.Header().Set("Content-Type", "application/json") + require.NoError(t, json.NewEncoder(w).Encode(response)) +} + +func testClientContext() ClientContext { + return ClientContext{ + UserID: "12345", + ProToken: "secret-token", + Locale: "fa-IR", + Platform: "macos", + AppVersion: "9.2.0", + } +} + +func testMessage(displayID string, expiresAt time.Time) *wire.ResolvedUserMessage { + return &wire.ResolvedUserMessage{ + DisplayID: displayID, + CampaignID: "campaign-1", + RevisionID: "revision-1", + DeliveryID: "delivery-1", + Surface: wire.SurfaceSnackbar, + Locale: "fa-IR", + Body: "A safe localized message", + ExpiresAt: expiresAt, + } +} diff --git a/usermessage/service.go b/usermessage/service.go new file mode 100644 index 00000000..ae2ab910 --- /dev/null +++ b/usermessage/service.go @@ -0,0 +1,417 @@ +package usermessage + +import ( + "context" + "errors" + "log/slog" + "math/rand" + "sync" + "time" + + wire "github.com/getlantern/common/usermessage" +) + +const ( + initialFailureBackoff = 5 * time.Second + maxFailureBackoff = 5 * time.Minute +) + +// Clock creates timers and reports the current time. +type Clock interface { + Now() time.Time + NewTimer(time.Duration) Timer +} + +// Timer is the subset of time.Timer used by Service. +type Timer interface { + C() <-chan time.Time + Stop() bool +} + +type realClock struct{} + +func (realClock) Now() time.Time { return time.Now() } +func (realClock) NewTimer(d time.Duration) Timer { return realTimer{time.NewTimer(d)} } + +type realTimer struct{ *time.Timer } + +func (t realTimer) C() <-chan time.Time { return t.Timer.C } + +// Options configures a Service. +type Options struct { + DataDir string + Fetcher Fetcher + ContextProvider func() ClientContext + Clock Clock + Jitter func(time.Duration) time.Duration + Logger *slog.Logger +} + +// Service owns polling, per-account presentation state, and display acknowledgment. +type Service struct { + fetcher Fetcher + contextProvider func() ClientContext + clock Clock + jitter func(time.Duration) time.Duration + logger *slog.Logger + store *store + wake chan struct{} + + mu sync.Mutex + started bool + active bool + online bool + generation uint64 + requestID uint64 + requestCancel context.CancelFunc +} + +// New creates a user-message service and loads its durable state. +func New(opts Options) (*Service, error) { + if opts.Fetcher == nil { + return nil, errors.New("user-message fetcher is required") + } + if opts.ContextProvider == nil { + return nil, errors.New("user-message context provider is required") + } + state, err := newStore(opts.DataDir) + if err != nil { + return nil, err + } + clock := opts.Clock + if clock == nil { + clock = realClock{} + } + jitter := opts.Jitter + if jitter == nil { + jitter = defaultJitter + } + logger := opts.Logger + if logger == nil { + logger = slog.Default().With("service", "user_messages") + } + return &Service{ + fetcher: opts.Fetcher, + contextProvider: opts.ContextProvider, + clock: clock, + jitter: jitter, + logger: logger, + store: state, + wake: make(chan struct{}, 1), + active: true, + online: true, + }, nil +} + +// Start begins with an immediate fetch and is idempotent. +func (s *Service) Start(ctx context.Context) { + s.mu.Lock() + if s.started { + s.mu.Unlock() + return + } + s.started = true + s.mu.Unlock() + go s.run(ctx) +} + +// Current returns the pending, unexpired message for the current account. +func (s *Service) Current() (*wire.ResolvedUserMessage, error) { + clientContext := s.contextProvider() + if clientContext.UserID == "" { + return nil, nil + } + return s.store.current(clientContext.UserID, s.clock.Now()) +} + +// Refresh requests an immediate fetch. Concurrent requests are coalesced. +func (s *Service) Refresh() { + s.mu.Lock() + s.generation++ + cancel := s.requestCancel + s.mu.Unlock() + if cancel != nil { + cancel() + } + s.signalRefresh() +} + +func (s *Service) signalRefresh() { + select { + case s.wake <- struct{}{}: + default: + } +} + +// Acknowledge marks a pending message as displayed and refreshes eligibility. +func (s *Service) Acknowledge(displayID string) error { + clientContext := s.contextProvider() + if clientContext.UserID == "" { + return ErrMessageNotPending + } + if err := s.store.acknowledge(clientContext.UserID, displayID, s.clock.Now()); err != nil { + return err + } + s.Refresh() + return nil +} + +// SetActivity controls polling while the host app is active and online. +func (s *Service) SetActivity(active, online bool) { + s.mu.Lock() + changed := s.active != active || s.online != online + s.active = active + s.online = online + s.mu.Unlock() + if changed { + s.Refresh() + } +} + +func (s *Service) run(ctx context.Context) { + var delay time.Duration + var failures uint + for s.wait(ctx, delay) { + requestContext, requestID, generation, ok := s.beginRequest(ctx) + if !ok { + delay = 0 + continue + } + clientContext := s.contextProvider() + if !clientContext.valid() { + s.endRequest(requestID) + failures = 0 + delay = 0 + s.logger.Debug("User-message fetch deferred", "reason", "credentials_unavailable") + if !s.waitForRefresh(ctx) { + return + } + continue + } + seen := s.store.seen(clientContext.UserID) + response, err := s.fetcher.Fetch(requestContext, clientContext, seen) + s.endRequest(requestID) + if errors.Is(err, context.Canceled) { + s.consumeRefresh() + delay = 0 + continue + } + if err == nil { + pollOnly := response + pollOnly.Message = nil + err = pollOnly.Validate() + if err != nil { + failures++ + delay = s.jitter(failureBackoff(failures)) + s.logger.Warn( + "User-message response rejected", + "category", "invalid_response", + "failure_count", failures, + "retry_in", delay, + ) + continue + } + if response.Message != nil && response.Message.Validate() != nil { + s.logger.Warn("User-message response discarded", "category", "invalid_message") + response.Message = nil + } + } + if err != nil { + failures++ + delay = s.jitter(failureBackoff(failures)) + s.logFetchFailure(err, failures, delay) + continue + } + if s.generationChanged(generation) || s.contextProvider() != clientContext { + s.consumeRefresh() + delay = 0 + continue + } + if err := s.store.offer(clientContext.UserID, response.Message, s.clock.Now()); err != nil { + failures++ + delay = s.jitter(failureBackoff(failures)) + s.logger.Warn( + "User-message fetch result could not be persisted", + "category", "local_state", + "failure_count", failures, + "retry_in", delay, + ) + continue + } + failures = 0 + delay = s.jitter(time.Duration(response.PollIntervalSeconds) * time.Second) + s.logFetchResult(response.Message, delay) + } +} + +func (s *Service) waitForRefresh(ctx context.Context) bool { + select { + case <-ctx.Done(): + return false + case <-s.wake: + return true + } +} + +func (s *Service) logFetchFailure(err error, failures uint, retryIn time.Duration) { + category, statusCode := fetchFailureDetails(err) + attributes := []any{ + "category", category, + "failure_count", failures, + "retry_in", retryIn, + } + if statusCode != 0 { + attributes = append(attributes, "http_status", statusCode) + } + // Do not attach err here. Transport errors can contain request URLs, and + // future error wrappers might include credentials or localized content. + s.logger.Warn("User-message fetch failed", attributes...) +} + +func fetchFailureDetails(err error) (category string, statusCode int) { + if errors.Is(err, errCredentialsUnavailable) { + return "credentials_unavailable", 0 + } + if errors.Is(err, context.DeadlineExceeded) { + return "timeout", 0 + } + var statusErr *httpStatusError + if !errors.As(err, &statusErr) { + return "transport", 0 + } + statusCode = statusErr.statusCode + switch { + case statusCode == 400: + category = "request_rejected" + case statusCode == 401 || statusCode == 403: + category = "authentication" + case statusCode == 404: + category = "endpoint" + case statusCode == 408: + category = "timeout" + case statusCode == 429: + category = "rate_limited" + case statusCode >= 500: + category = "server" + default: + category = "http" + } + return category, statusCode +} + +func (s *Service) logFetchResult(message *wire.ResolvedUserMessage, pollIn time.Duration) { + if message == nil { + s.logger.Debug("User-message fetch completed", "result", "no_message", "poll_in", pollIn) + return + } + s.logger.Info( + "User-message fetch completed", + "result", "message_available", + "campaign_id", message.CampaignID, + "revision_id", message.RevisionID, + "delivery_id", message.DeliveryID, + "surface", message.Surface, + "locale", message.Locale, + "expires_at", message.ExpiresAt, + "poll_in", pollIn, + ) +} + +func (s *Service) beginRequest(parent context.Context) (context.Context, uint64, uint64, bool) { + requestContext, cancel := context.WithCancel(parent) + s.mu.Lock() + defer s.mu.Unlock() + if !s.active || !s.online { + cancel() + return nil, 0, 0, false + } + s.requestID++ + s.requestCancel = cancel + return requestContext, s.requestID, s.generation, true +} + +func (s *Service) endRequest(requestID uint64) { + s.mu.Lock() + var cancel context.CancelFunc + if s.requestID == requestID { + cancel = s.requestCancel + s.requestCancel = nil + } + s.mu.Unlock() + if cancel != nil { + cancel() + } +} + +func (s *Service) generationChanged(generation uint64) bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.generation != generation +} + +func (s *Service) consumeRefresh() { + select { + case <-s.wake: + default: + } +} + +func (s *Service) wait(ctx context.Context, delay time.Duration) bool { + for { + if ctx.Err() != nil { + return false + } + if !s.ready() { + select { + case <-ctx.Done(): + return false + case <-s.wake: + if s.ready() { + return true + } + continue + } + } + if delay <= 0 { + return true + } + timer := s.clock.NewTimer(delay) + select { + case <-ctx.Done(): + timer.Stop() + return false + case <-s.wake: + timer.Stop() + if s.ready() { + return true + } + continue + case <-timer.C(): + return true + } + } +} + +func (s *Service) ready() bool { + s.mu.Lock() + defer s.mu.Unlock() + return s.active && s.online +} + +func failureBackoff(failures uint) time.Duration { + delay := initialFailureBackoff + for i := uint(1); i < failures && delay < maxFailureBackoff; i++ { + delay *= 2 + if delay >= maxFailureBackoff { + return maxFailureBackoff + } + } + return delay +} + +func defaultJitter(delay time.Duration) time.Duration { + if delay <= 0 { + return 0 + } + return time.Duration(float64(delay) * (0.9 + rand.Float64()*0.1)) +} diff --git a/usermessage/service_test.go b/usermessage/service_test.go new file mode 100644 index 00000000..59503b66 --- /dev/null +++ b/usermessage/service_test.go @@ -0,0 +1,415 @@ +package usermessage + +import ( + "bytes" + "context" + "errors" + "io" + "log/slog" + "net/http" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" + + wire "github.com/getlantern/common/usermessage" +) + +func TestServicePollingBackoffAndSuccessReset(t *testing.T) { + clock := newFakeClock(time.Date(2026, 8, 18, 12, 0, 0, 0, time.UTC)) + fetcher := newScriptedFetcher() + service := newTestService(t, clock, fetcher, func() ClientContext { return testClientContext() }) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + service.Start(ctx) + + receiveFetch(t, fetcher) + fetcher.results <- fetchResult{err: errors.New("offline")} + require.Equal(t, 5*time.Second, receiveTimer(t, clock)) + clock.Advance(5 * time.Second) + + receiveFetch(t, fetcher) + fetcher.results <- fetchResult{err: errors.New("still offline")} + require.Equal(t, 10*time.Second, receiveTimer(t, clock)) + clock.Advance(10 * time.Second) + + receiveFetch(t, fetcher) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{PollIntervalSeconds: 300}} + require.Equal(t, 5*time.Minute, receiveTimer(t, clock)) + clock.Advance(5 * time.Minute) + + receiveFetch(t, fetcher) + fetcher.results <- fetchResult{err: errors.New("failed after success")} + require.Equal(t, 5*time.Second, receiveTimer(t, clock)) +} + +func TestServiceImmediateRefreshForContextAndActivityChanges(t *testing.T) { + clock := newFakeClock(time.Now()) + fetcher := newScriptedFetcher() + var mu sync.Mutex + clientContext := testClientContext() + provider := func() ClientContext { + mu.Lock() + defer mu.Unlock() + return clientContext + } + service := newTestService(t, clock, fetcher, provider) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + service.Start(ctx) + + require.Equal(t, "fa-IR", receiveFetch(t, fetcher).clientContext.Locale) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{PollIntervalSeconds: 300}} + receiveTimer(t, clock) + + mu.Lock() + clientContext.Locale = "en-US" + mu.Unlock() + service.Refresh() + require.Equal(t, "en-US", receiveFetch(t, fetcher).clientContext.Locale) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{PollIntervalSeconds: 300}} + receiveTimer(t, clock) + + service.SetActivity(false, true) + require.Never(t, func() bool { + select { + case <-fetcher.requests: + return true + default: + return false + } + }, 50*time.Millisecond, time.Millisecond) + service.SetActivity(true, true) + require.Equal(t, "en-US", receiveFetch(t, fetcher).clientContext.Locale) +} + +func TestServiceSeenFilteringAccountSwitchAndExpiration(t *testing.T) { + clock := newFakeClock(time.Now()) + fetcher := newScriptedFetcher() + var mu sync.Mutex + clientContext := testClientContext() + provider := func() ClientContext { + mu.Lock() + defer mu.Unlock() + return clientContext + } + service := newTestService(t, clock, fetcher, provider) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + service.Start(ctx) + + first := receiveFetch(t, fetcher) + require.Empty(t, first.seen) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{ + PollIntervalSeconds: 300, + Message: testMessage("display-1", clock.Now().Add(time.Minute)), + }} + receiveTimer(t, clock) + message, err := service.Current() + require.NoError(t, err) + require.Equal(t, "display-1", message.DisplayID) + require.NoError(t, service.Acknowledge("display-1")) + request := receiveFetch(t, fetcher) + require.Equal(t, []string{"display-1"}, request.seen) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{PollIntervalSeconds: 300}} + receiveTimer(t, clock) + + mu.Lock() + clientContext.UserID = "67890" + mu.Unlock() + service.Refresh() + request = receiveFetch(t, fetcher) + require.Empty(t, request.seen) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{ + PollIntervalSeconds: 300, + Message: testMessage("display-2", clock.Now().Add(time.Minute)), + }} + receiveTimer(t, clock) + clock.Advance(time.Minute) + message, err = service.Current() + require.NoError(t, err) + require.Nil(t, message) +} + +func TestServiceCancelsFetchAfterAccountReplacement(t *testing.T) { + clock := newFakeClock(time.Now()) + fetcher := newScriptedFetcher() + var mu sync.Mutex + clientContext := testClientContext() + provider := func() ClientContext { + mu.Lock() + defer mu.Unlock() + return clientContext + } + service := newTestService(t, clock, fetcher, provider) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + service.Start(ctx) + + require.Equal(t, "12345", receiveFetch(t, fetcher).clientContext.UserID) + mu.Lock() + clientContext.UserID = "67890" + mu.Unlock() + service.Refresh() + + require.Equal(t, "67890", receiveFetch(t, fetcher).clientContext.UserID) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{PollIntervalSeconds: 300}} + receiveTimer(t, clock) + require.Never(t, func() bool { + select { + case <-fetcher.requests: + return true + default: + return false + } + }, 50*time.Millisecond, time.Millisecond) + + mu.Lock() + clientContext.UserID = "12345" + mu.Unlock() + message, err := service.Current() + require.NoError(t, err) + require.Nil(t, message) +} + +func TestServiceWaitsForCompleteCredentials(t *testing.T) { + clock := newFakeClock(time.Now()) + fetcher := newScriptedFetcher() + var mu sync.Mutex + clientContext := ClientContext{} + provider := func() ClientContext { + mu.Lock() + defer mu.Unlock() + return clientContext + } + service := newTestService(t, clock, fetcher, provider) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + service.Start(ctx) + + require.Never(t, func() bool { + select { + case <-fetcher.requests: + return true + default: + return false + } + }, 50*time.Millisecond, time.Millisecond) + + mu.Lock() + clientContext = testClientContext() + mu.Unlock() + service.Refresh() + require.Equal(t, "12345", receiveFetch(t, fetcher).clientContext.UserID) +} + +func TestServiceLogsSafeFetchOutcomes(t *testing.T) { + clock := newFakeClock(time.Now()) + fetcher := newScriptedFetcher() + var logs bytes.Buffer + logger := slog.New(slog.NewTextHandler(&logs, &slog.HandlerOptions{Level: slog.LevelDebug})) + service, err := New(Options{ + DataDir: t.TempDir(), + Fetcher: fetcher, + ContextProvider: func() ClientContext { return testClientContext() }, + Clock: clock, + Jitter: func(delay time.Duration) time.Duration { return delay }, + Logger: logger, + }) + require.NoError(t, err) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + service.Start(ctx) + + receiveFetch(t, fetcher) + fetcher.results <- fetchResult{err: &httpStatusError{statusCode: http.StatusUnauthorized}} + require.Equal(t, 5*time.Second, receiveTimer(t, clock)) + require.Contains(t, logs.String(), "category=authentication") + require.Contains(t, logs.String(), "http_status=401") + + clock.Advance(5 * time.Second) + receiveFetch(t, fetcher) + fetcher.results <- fetchResult{response: wire.UserMessageResponse{ + PollIntervalSeconds: 300, + Message: testMessage("display-1", clock.Now().Add(time.Hour)), + }} + require.Equal(t, 5*time.Minute, receiveTimer(t, clock)) + require.Contains(t, logs.String(), "result=message_available") + require.NotContains(t, logs.String(), "12345") + require.NotContains(t, logs.String(), "secret-token") + require.NotContains(t, logs.String(), "A safe localized message") +} + +func TestServiceStopsAfterParentContextCancellation(t *testing.T) { + clock := newFakeClock(time.Now()) + fetcher := newScriptedFetcher() + service := newTestService(t, clock, fetcher, func() ClientContext { return testClientContext() }) + ctx, cancel := context.WithCancel(context.Background()) + service.Start(ctx) + receiveFetch(t, fetcher) + cancel() + + require.Never(t, func() bool { + select { + case <-fetcher.requests: + return true + default: + return false + } + }, 50*time.Millisecond, time.Millisecond) +} + +func TestContextNormalization(t *testing.T) { + require.Equal(t, "macos", NormalizePlatform("darwin")) + require.Equal(t, "windows", NormalizePlatform(" Windows ")) + require.Equal(t, "fa-IR", NormalizeLocale("fa-ir")) + require.Equal(t, "en-US", NormalizeLocale("not a locale")) +} + +func TestPollingDelaysStayWithinSLAAndBackoffCap(t *testing.T) { + for range 100 { + delay := defaultJitter(5 * time.Minute) + require.GreaterOrEqual(t, delay, 270*time.Second) + require.LessOrEqual(t, delay, 5*time.Minute) + } + require.Equal(t, 5*time.Minute, failureBackoff(100)) +} + +type fetchCall struct { + clientContext ClientContext + seen []string +} + +type fetchResult struct { + response wire.UserMessageResponse + err error +} + +type scriptedFetcher struct { + requests chan fetchCall + results chan fetchResult +} + +func newScriptedFetcher() *scriptedFetcher { + return &scriptedFetcher{ + requests: make(chan fetchCall, 8), + results: make(chan fetchResult, 8), + } +} + +func (f *scriptedFetcher) Fetch( + ctx context.Context, + clientContext ClientContext, + seen []string, +) (wire.UserMessageResponse, error) { + select { + case f.requests <- fetchCall{clientContext: clientContext, seen: seen}: + case <-ctx.Done(): + return wire.UserMessageResponse{}, ctx.Err() + } + select { + case result := <-f.results: + return result.response, result.err + case <-ctx.Done(): + return wire.UserMessageResponse{}, ctx.Err() + } +} + +func newTestService( + t *testing.T, + clock Clock, + fetcher Fetcher, + provider func() ClientContext, +) *Service { + t.Helper() + service, err := New(Options{ + DataDir: t.TempDir(), + Fetcher: fetcher, + ContextProvider: provider, + Clock: clock, + Jitter: func(delay time.Duration) time.Duration { return delay }, + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + }) + require.NoError(t, err) + return service +} + +func receiveFetch(t *testing.T, fetcher *scriptedFetcher) fetchCall { + t.Helper() + select { + case request := <-fetcher.requests: + return request + case <-time.After(time.Second): + t.Fatal("timed out waiting for fetch") + return fetchCall{} + } +} + +func receiveTimer(t *testing.T, clock *fakeClock) time.Duration { + t.Helper() + select { + case delay := <-clock.created: + return delay + case <-time.After(time.Second): + t.Fatal("timed out waiting for timer") + return 0 + } +} + +type fakeClock struct { + mu sync.Mutex + now time.Time + timers []*fakeTimer + created chan time.Duration +} + +func newFakeClock(now time.Time) *fakeClock { + return &fakeClock{now: now, created: make(chan time.Duration, 16)} +} + +func (c *fakeClock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.now +} + +func (c *fakeClock) NewTimer(delay time.Duration) Timer { + c.mu.Lock() + defer c.mu.Unlock() + timer := &fakeTimer{clock: c, due: c.now.Add(delay), ch: make(chan time.Time, 1)} + c.timers = append(c.timers, timer) + c.created <- delay + return timer +} + +func (c *fakeClock) Advance(delay time.Duration) { + c.mu.Lock() + c.now = c.now.Add(delay) + now := c.now + for _, timer := range c.timers { + if !timer.stopped && !timer.fired && !timer.due.After(now) { + timer.fired = true + timer.ch <- now + } + } + c.mu.Unlock() +} + +type fakeTimer struct { + clock *fakeClock + due time.Time + ch chan time.Time + stopped bool + fired bool +} + +func (t *fakeTimer) C() <-chan time.Time { return t.ch } + +func (t *fakeTimer) Stop() bool { + t.clock.mu.Lock() + defer t.clock.mu.Unlock() + wasActive := !t.stopped && !t.fired + t.stopped = true + return wasActive +} diff --git a/usermessage/store.go b/usermessage/store.go new file mode 100644 index 00000000..d17062a4 --- /dev/null +++ b/usermessage/store.go @@ -0,0 +1,271 @@ +package usermessage + +import ( + "encoding/json" + "errors" + "fmt" + "os" + "path/filepath" + "slices" + "strings" + "sync" + "time" + + wire "github.com/getlantern/common/usermessage" + + "github.com/getlantern/radiance/common/atomicfile" + "github.com/getlantern/radiance/common/fileperm" +) + +const ( + stateVersion = 1 + maxUsers = 16 +) + +// ErrMessageNotPending indicates that a display ID cannot be acknowledged for the current account. +var ErrMessageNotPending = errors.New("user message is not pending") + +type persistedState struct { + Version int `json:"version"` + Users map[string]*userState `json:"users,omitempty"` + Order []string `json:"order,omitempty"` +} + +type userState struct { + Seen []string `json:"seen,omitempty"` + Pending *wire.ResolvedUserMessage `json:"pending,omitempty"` +} + +type store struct { + mu sync.Mutex + path string + state persistedState +} + +func newStore(dataDir string) (*store, error) { + s := &store{ + path: filepath.Join(dataDir, "user-messages.json"), + state: persistedState{ + Version: stateVersion, + Users: make(map[string]*userState), + }, + } + data, err := os.ReadFile(s.path) + if errors.Is(err, os.ErrNotExist) { + return s, nil + } + if err != nil { + return nil, fmt.Errorf("read user-message state: %w", err) + } + if err := json.Unmarshal(data, &s.state); err != nil { + return nil, fmt.Errorf("decode user-message state: %w", err) + } + if s.state.Version != stateVersion { + return nil, fmt.Errorf("unsupported user-message state version %d", s.state.Version) + } + if s.state.Users == nil { + s.state.Users = make(map[string]*userState) + } + s.sanitize() + return s, nil +} + +func (s *store) seen(userID string) []string { + s.mu.Lock() + defer s.mu.Unlock() + state := s.state.Users[userID] + if state == nil { + return nil + } + return slices.Clone(state.Seen) +} + +func (s *store) current(userID string, now time.Time) (*wire.ResolvedUserMessage, error) { + s.mu.Lock() + defer s.mu.Unlock() + state := s.state.Users[userID] + if state == nil || state.Pending == nil { + return nil, nil + } + if !now.Before(state.Pending.ExpiresAt) { + next := cloneState(s.state) + next.Users[userID].Pending = nil + if err := s.commitLocked(next); err != nil { + return nil, err + } + return nil, nil + } + return cloneMessage(state.Pending), nil +} + +func (s *store) offer(userID string, message *wire.ResolvedUserMessage, now time.Time) error { + s.mu.Lock() + defer s.mu.Unlock() + next := cloneState(s.state) + state := next.Users[userID] + expired := state != nil && state.Pending != nil && !now.Before(state.Pending.ExpiresAt) + if expired { + state.Pending = nil + } + if message == nil || !now.Before(message.ExpiresAt) { + if expired { + return s.commitLocked(next) + } + return nil + } + if state == nil { + state = &userState{} + next.Users[userID] = state + touch(&next, userID) + } + if slices.Contains(state.Seen, message.DisplayID) || state.Pending != nil { + return nil + } + state.Pending = cloneMessage(message) + return s.commitLocked(next) +} + +func (s *store) acknowledge(userID, displayID string, now time.Time) error { + s.mu.Lock() + defer s.mu.Unlock() + state := s.state.Users[userID] + if state == nil { + return ErrMessageNotPending + } + if slices.Contains(state.Seen, displayID) { + return nil + } + if state.Pending == nil || state.Pending.DisplayID != displayID || !now.Before(state.Pending.ExpiresAt) { + if state.Pending != nil && !now.Before(state.Pending.ExpiresAt) { + next := cloneState(s.state) + next.Users[userID].Pending = nil + if err := s.commitLocked(next); err != nil { + return err + } + } + return ErrMessageNotPending + } + next := cloneState(s.state) + state = next.Users[userID] + state.Pending = nil + state.Seen = append(state.Seen, displayID) + if len(state.Seen) > wire.MaxSeenDisplayIDs { + state.Seen = slices.Clone(state.Seen[len(state.Seen)-wire.MaxSeenDisplayIDs:]) + } + touch(&next, userID) + return s.commitLocked(next) +} + +func touch(state *persistedState, userID string) { + state.Order = slices.DeleteFunc(state.Order, func(id string) bool { return id == userID }) + state.Order = append(state.Order, userID) + for len(state.Order) > maxUsers { + delete(state.Users, state.Order[0]) + state.Order = state.Order[1:] + } +} + +func (s *store) saveLocked() error { + return writeState(s.path, s.state) +} + +func (s *store) commitLocked(next persistedState) error { + if err := writeState(s.path, next); err != nil { + return err + } + s.state = next + return nil +} + +func writeState(path string, state persistedState) error { + data, err := json.Marshal(state) + if err != nil { + return fmt.Errorf("encode user-message state: %w", err) + } + if err := atomicfile.WriteFile(path, data, fileperm.File); err != nil { + return fmt.Errorf("write user-message state: %w", err) + } + return nil +} + +func cloneState(state persistedState) persistedState { + clone := persistedState{ + Version: state.Version, + Users: make(map[string]*userState, len(state.Users)), + Order: slices.Clone(state.Order), + } + for userID, current := range state.Users { + clone.Users[userID] = &userState{ + Seen: slices.Clone(current.Seen), + Pending: cloneMessage(current.Pending), + } + } + return clone +} + +func (s *store) sanitize() { + validUsers := make(map[string]*userState, len(s.state.Users)) + for userID, state := range s.state.Users { + if userID == "" || state == nil { + continue + } + seen := make([]string, 0, min(len(state.Seen), wire.MaxSeenDisplayIDs)) + for _, id := range state.Seen { + if validDisplayID(id) && !slices.Contains(seen, id) { + seen = append(seen, id) + } + } + if len(seen) > wire.MaxSeenDisplayIDs { + seen = seen[len(seen)-wire.MaxSeenDisplayIDs:] + } + state.Seen = seen + if state.Pending != nil && state.Pending.Validate() != nil { + state.Pending = nil + } + validUsers[userID] = state + } + s.state.Users = validUsers + order := make([]string, 0, min(len(s.state.Order), maxUsers)) + for _, userID := range s.state.Order { + if _, ok := validUsers[userID]; ok && !slices.Contains(order, userID) { + order = append(order, userID) + } + } + for userID := range validUsers { + if !slices.Contains(order, userID) { + order = append(order, userID) + } + } + if len(order) > maxUsers { + for _, userID := range order[:len(order)-maxUsers] { + delete(validUsers, userID) + } + order = order[len(order)-maxUsers:] + } + s.state.Order = order +} + +func validDisplayID(id string) bool { + if id == "" || len(id) > wire.MaxDisplayIDLength { + return false + } + for _, r := range id { + if r > 127 || !(r >= 'a' && r <= 'z') && !(r >= 'A' && r <= 'Z') && + !(r >= '0' && r <= '9') && !strings.ContainsRune("._:-", r) { + return false + } + } + return true +} + +func cloneMessage(message *wire.ResolvedUserMessage) *wire.ResolvedUserMessage { + if message == nil { + return nil + } + clone := *message + if message.Action != nil { + action := *message.Action + clone.Action = &action + } + return &clone +} diff --git a/usermessage/store_test.go b/usermessage/store_test.go new file mode 100644 index 00000000..f89d9fd7 --- /dev/null +++ b/usermessage/store_test.go @@ -0,0 +1,100 @@ +package usermessage + +import ( + "errors" + "fmt" + "path/filepath" + "testing" + "time" + + "github.com/stretchr/testify/require" + + wire "github.com/getlantern/common/usermessage" +) + +func TestStorePersistsPendingAndSeenByUser(t *testing.T) { + now := time.Date(2026, 8, 18, 12, 0, 0, 0, time.UTC) + dir := t.TempDir() + state, err := newStore(dir) + require.NoError(t, err) + require.NoError(t, state.offer("1", testMessage("display-1", now.Add(time.Hour)), now)) + require.NoError(t, state.offer("2", testMessage("display-2", now.Add(time.Hour)), now)) + + reloaded, err := newStore(dir) + require.NoError(t, err) + message, err := reloaded.current("1", now) + require.NoError(t, err) + require.Equal(t, "display-1", message.DisplayID) + require.NoError(t, reloaded.acknowledge("1", "display-1", now)) + require.NoError(t, reloaded.acknowledge("1", "display-1", now)) + require.Equal(t, []string{"display-1"}, reloaded.seen("1")) + require.Empty(t, reloaded.seen("2")) + + reloadedAgain, err := newStore(dir) + require.NoError(t, err) + require.Equal(t, []string{"display-1"}, reloadedAgain.seen("1")) + message, err = reloadedAgain.current("1", now) + require.NoError(t, err) + require.Nil(t, message) +} + +func TestStoreBoundsSeenIDsAndExpiresPending(t *testing.T) { + now := time.Date(2026, 8, 18, 12, 0, 0, 0, time.UTC) + state, err := newStore(t.TempDir()) + require.NoError(t, err) + for i := 0; i < wire.MaxSeenDisplayIDs+3; i++ { + id := fmt.Sprintf("display-%d", i) + require.NoError(t, state.offer("1", testMessage(id, now.Add(time.Hour)), now)) + require.NoError(t, state.acknowledge("1", id, now)) + } + seen := state.seen("1") + require.Len(t, seen, wire.MaxSeenDisplayIDs) + require.Equal(t, "display-3", seen[0]) + + require.NoError(t, state.offer("1", testMessage("expiring", now.Add(time.Minute)), now)) + message, err := state.current("1", now.Add(time.Minute)) + require.NoError(t, err) + require.Nil(t, message) + require.ErrorIs(t, state.acknowledge("1", "expiring", now.Add(time.Minute)), ErrMessageNotPending) +} + +func TestStoreDoesNotReplaceUnacknowledgedMessage(t *testing.T) { + now := time.Now() + state, err := newStore(t.TempDir()) + require.NoError(t, err) + require.NoError(t, state.offer("1", testMessage("first", now.Add(time.Hour)), now)) + require.NoError(t, state.offer("1", testMessage("second", now.Add(time.Hour)), now)) + message, err := state.current("1", now) + require.NoError(t, err) + require.Equal(t, "first", message.DisplayID) + require.True(t, errors.Is(state.acknowledge("1", "second", now), ErrMessageNotPending)) +} + +func TestStoreSanitizesInvalidPersistedSeenIDs(t *testing.T) { + state, err := newStore(t.TempDir()) + require.NoError(t, err) + state.state.Users["1"] = &userState{Seen: []string{"valid-id", "invalid id", "valid-id"}} + state.state.Order = []string{"1"} + require.NoError(t, state.saveLocked()) + + reloaded, err := newStore(filepath.Dir(state.path)) + require.NoError(t, err) + require.Equal(t, []string{"valid-id"}, reloaded.seen("1")) +} + +func TestStoreKeepsPendingWhenAcknowledgmentWriteFails(t *testing.T) { + now := time.Now() + dir := t.TempDir() + state, err := newStore(dir) + require.NoError(t, err) + require.NoError(t, state.offer("1", testMessage("display-1", now.Add(time.Hour)), now)) + + validPath := state.path + state.path = dir + require.Error(t, state.acknowledge("1", "display-1", now)) + state.path = validPath + message, err := state.current("1", now) + require.NoError(t, err) + require.Equal(t, "display-1", message.DisplayID) + require.Empty(t, state.seen("1")) +}