diff --git a/cmd/soundtouch-service/main.go b/cmd/soundtouch-service/main.go index 0af013d..3ca38d3 100644 --- a/cmd/soundtouch-service/main.go +++ b/cmd/soundtouch-service/main.go @@ -216,6 +216,12 @@ func main() { Usage: "Log what would be migrated without actually doing it", EnvVars: []string{"MIGRATION_DRY_RUN"}, }, + &cli.StringFlag{ + Name: "preferred-source", + Usage: "Preferred source of truth (local or upstream)", + Value: "local", + EnvVars: []string{"PREFERRED_SOURCE"}, + }, }, Action: func(c *cli.Context) error { config := loadConfig(c) @@ -247,7 +253,7 @@ func main() { server.SetVersionInfo(version, commit, date) server.SetDiscoverySettings(config.discoveryInterval, persisted.DiscoveryEnabled) server.SetDNSSettings(persisted.DNSEnabled, strings.Join(persisted.DNSUpstream, ","), persisted.DNSBindAddr) - server.SetMirrorSettings(persisted.MirrorEnabled, persisted.MirrorEndpoints) + server.SetMirrorSettings(persisted.MirrorEnabled, persisted.MirrorEndpoints, persisted.PreferredSource) server.SetInternalPaths(persisted.InternalPaths) server.SetSpotifyConfig(config.spotifyClientID, config.spotifyClientSecret, config.spotifyRedirectURI) server.SetMgmtConfig(config.mgmtUsername, config.mgmtPassword) @@ -392,6 +398,7 @@ type serviceConfig struct { mgmtPassword string migrationEnabled bool migrationDryRun bool + preferredSource string } func loadConfig(c *cli.Context) serviceConfig { @@ -460,6 +467,7 @@ func loadConfig(c *cli.Context) serviceConfig { internalPaths := c.StringSlice("internal-paths") migrationEnabled := c.Bool("migration-enabled") migrationDryRun := c.Bool("migration-dry-run") + preferredSource := c.String("preferred-source") return serviceConfig{ port: port, @@ -489,6 +497,7 @@ func loadConfig(c *cli.Context) serviceConfig { mgmtPassword: mgmtPassword, migrationEnabled: migrationEnabled, migrationDryRun: migrationDryRun, + preferredSource: preferredSource, } } @@ -498,14 +507,17 @@ func getDomains(serverURL, httpsServerURL, hostname string) []string { "*.api.bose.io": true, "*.api.bosecm.com": true, // Core Bose domains (keep specific ones for clarity) - "streaming.bose.com": true, - "updates.bose.com": true, - "stats.bose.com": true, - "bmx.bose.com": true, - "worldwide.bose.com": true, - "music.api.bose.com": true, - "bose-prod.apigee.net": true, - "bose-test.apigee.net": true, + "streaming.bose.com": true, + "updates.bose.com": true, + "stats.bose.com": true, + "bmx.bose.com": true, + "worldwide.bose.com": true, + "music.api.bose.com": true, + "streamingoauth.bose.com": true, + "bosecm.com": true, + "bose.io": true, + "bose-prod.apigee.net": true, + "bose-test.apigee.net": true, // Local service domains setup.TestDomain: true, hostname: true, @@ -535,6 +547,13 @@ func applyPersistedSettings(ds *datastore.DataStore, config *serviceConfig) data return datastore.Settings{} } + // Only override CLI values if settings file exists + // If no settings file exists, GetSettings returns empty Settings{} and we should preserve CLI values + settingsPath := filepath.Join(ds.DataDir, "settings.json") + if _, err := os.Stat(settingsPath); os.IsNotExist(err) { + return datastore.Settings{} + } + if persisted.ServerURL != "" { config.serverURL = persisted.ServerURL } @@ -569,6 +588,7 @@ func applyPersistedSettings(ds *datastore.DataStore, config *serviceConfig) data config.mirrorEnabled = persisted.MirrorEnabled config.mirrorEndpoints = persisted.MirrorEndpoints + config.preferredSource = persisted.PreferredSource config.internalPaths = persisted.InternalPaths return persisted @@ -590,12 +610,14 @@ func createDefaultSettings(ds *datastore.DataStore, config serviceConfig) datast DNSBindAddr: config.dnsBind, MirrorEnabled: config.mirrorEnabled, MirrorEndpoints: config.mirrorEndpoints, + PreferredSource: config.preferredSource, InternalPaths: config.internalPaths, Shortcuts: map[string]int{ "/.well-known/appspecific/com.chrome.devtools.json": http.StatusNotFound, "/sw.js": http.StatusNotFound, }, } + _ = ds.SaveSettings(settings) return settings @@ -634,6 +656,7 @@ func startDeviceDiscovery(server *handlers.Server) { func setupRouter(server *handlers.Server) *chi.Mux { r := chi.NewRouter() + r.Use(server.SnapshotMiddleware) r.Use(server.OriginMiddleware) r.Use(middleware.Recoverer) r.Use(server.ShortcutMiddleware) @@ -724,6 +747,10 @@ func setupRouter(server *handlers.Server) *chi.Mux { r.Post("/account/{account}/password", server.HandleMargeChangePassword) }) + r.Route("/oauth", func(r chi.Router) { + r.HandleFunc("/*", server.HandleBoseProxy) + }) + r.Route("/v1", func(r chi.Router) { r.Post("/stapp/{deviceId}", server.HandleAppEvents) r.Post("/scmudc/{deviceId}", server.HandleAppEvents) diff --git a/docs/REQUEST_RECORDING_CONCEPT.md b/docs/REQUEST_RECORDING_CONCEPT.md new file mode 100644 index 0000000..c5bfa4f --- /dev/null +++ b/docs/REQUEST_RECORDING_CONCEPT.md @@ -0,0 +1,337 @@ +# Request Recording Concept + +## Problem Statement + +The current request recording system has fundamental issues when dealing with request cloning, body consumption, and multiple response scenarios. Specifically: + +1. **Body Consumption**: HTTP request bodies can only be read once, leading to missing bodies in recordings +2. **Request Cloning**: A single original request may be cloned multiple times for different purposes (local handling, mirroring, recording) +3. **Multiple Responses**: The same logical request may generate different responses (local vs upstream mirror) +4. **Data Integrity**: No guarantee that recorded requests are identical across different execution paths + +## Current Issues (Examples) + +### Issue 1: Missing Request Bodies in Mirror Recordings + +**Local Recording** (complete): +```http +### POST /v1/scmudc/A81B6A536A98 +POST /v1/scmudc/A81B6A536A98 +Host: events.api.bosecm.com +Content-Type: text/json; charset=utf-8 +Content-Length: 587 +Authorization: Bearer jGwEmFWr... + +{"envelope":{"monoTime":234906,"payloadProtocolVersion":"3.1","payloadType":"scmudc","protocolVersion":"1.0","time":"2026-02-25T23:03:14.976349+00:00","uniqueId":"A81B6A536A98"},"payload":{"deviceInfo":{"boseID":"3230304","deviceID":"A81B6A536A98","deviceType":"SoundTouch 10","serialNumber":"I6332527703739342000020","softwareVersion":"27.0.6.46330.5043500 epdbuild.trunk.hepdswbld04.2022-08-04T11:20:29","systemSerialNumber":"069231P63364828AE"},"events":[{"data":{"play-state":"PAUSE_STATE"},"monoTime":234904,"time":"2026-02-25T23:03:14.973466+00:00","type":"play-state-changed"}]}} + +> {% + // Response: 200 OK +%} +``` + +**Mirror Recording** (missing body): +```http +### POST /v1/scmudc/A81B6A536A98 +POST /v1/scmudc/A81B6A536A98 +Host: events.api.bosecm.com +Content-Type: text/json; charset=utf-8 +Content-Length: 587 +Authorization: Bearer jGwEmFWr... + + + +> {% + // Response: 200 OK + // Headers: + // X-Proxy-Origin: upstream-mirror +%} +``` + +### Issue 2: Request Flow Complexity + +Current middleware execution order: +``` +1. MirrorMiddleware - Buffers body, creates clones +2. RecordMiddleware - Also buffers body +3. Application Handler - Processes request +4. Mirror Execution - Async/sync mirror to upstream +5. Recording - Multiple recording points +``` + +Problems: +- Multiple body reads across middleware chain +- Inconsistent request state between clones +- Race conditions in async scenarios +- No guarantee of request equivalence + +## Proposed Solution: Context-Bound Request Snapshots + +### Core Concept + +Create **immutable request snapshots** early in the request lifecycle and propagate them through the **Request Context**. This ensures all downstream consumers (Mirroring, Recording, Parity Check) use identical data without re-reading the request body. + +### Architecture (Context-Only) + +``` +┌─────────────────┐ +│ Original Request│ +└─────────┬───────┘ + │ + ▼ +┌─────────────────┐ ┌──────────────────┐ +│ Snapshot Creator│───▶│ Request Context │ +│ (Middleware) │ │ (Pointer-based) │ +└─────────┬───────┘ └──────────────────┘ + │ │ + ▼ │ (Safe for async) +┌─────────────────┐ │ +│ Middleware │◀─────────────┘ +│ Chain │ +└─────────┬───────┘ + │ + ┌───▼────┐ ┌─────────┐ ┌──────────────┐ + │ Local │ │ Mirror │ │ Recording │ + │Handler │ │Execution│ │ System │ + └────────┘ └─────────┘ └──────────────┘ +``` + +### Request Snapshot Structure + +```go +type RequestSnapshot struct { + Method string + URL *url.URL + Headers http.Header + Body []byte + Host string + Timestamp time.Time +} + +// Typed key for context safety +type contextKey struct{ name string } +var SnapshotKey = &contextKey{"request_snapshot"} +``` + +### Implementation Strategy + +#### Phase 1: Snapshot Middleware + +```go +func (s *Server) SnapshotMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // 1. Capture body once with size limit (e.g. 2MB) + body, _ := io.ReadAll(io.LimitReader(r.Body, 2*1024*1024)) + r.Body.Close() + + // 2. Create snapshot + snapshot := &RequestSnapshot{ + Method: r.Method, + URL: cloneURL(r.URL), + Headers: r.Header.Clone(), + Body: body, + Host: r.Host, + Timestamp: time.Now(), + } + + // 3. Inject pointer into context + ctx := context.WithValue(r.Context(), SnapshotKey, snapshot) + + // 4. Restore r.Body for downstream compatibility + r = r.WithContext(ctx) + r.Body = io.NopCloser(bytes.NewReader(snapshot.Body)) + + next.ServeHTTP(w, r) + }) +} +``` + +#### Phase 2: Downstream Consumption + +Consumers (Mirror/Record) retrieve the snapshot directly from context: + +```go +snapshot, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot) +if ok { + // Use snapshot.Body directly instead of io.ReadAll(r.Body) +} +``` + +## Hardware Considerations (Raspberry Pi Zero 2W) + +To protect MicroSD health and optimize for limited memory: + +1. **No Intermediate Disk Storage**: Snapshots exist only in memory; they are never written to disk until the final `.http` recording is generated. +2. **Memory Management**: Use `sync.Pool` for temporary buffers to reduce GC churn on the single-core/low-memory SoC. +3. **Automatic Cleanup**: Snapshots are naturally garbage collected once the Request Context and all child goroutines (detached mirrors/recordings) finish. +4. **Body Capping**: Strict limits on snapshot size prevent OOM (Out-of-Memory) conditions. + +#### Phase 2: Response Capture System + +```go +type ResponseRecorder struct { + http.ResponseWriter + snapshot *ResponseSnapshot + snapshotID string + source string + startTime time.Time +} + +func (r *ResponseRecorder) WriteHeader(statusCode int) { + r.snapshot.StatusCode = statusCode + r.snapshot.Headers = r.Header().Clone() + r.ResponseWriter.WriteHeader(statusCode) +} + +func (r *ResponseRecorder) Write(data []byte) (int, error) { + r.snapshot.Body = append(r.snapshot.Body, data...) + return r.ResponseWriter.Write(data) +} + +func (r *ResponseRecorder) finalize() { + r.snapshot.Duration = time.Since(r.startTime) + r.snapshot.Timestamp = time.Now() +} +``` + +#### Phase 3: Recording System Integration + +```go +type RecordingManager struct { + storage SnapshotStorage + recorder *Recorder + patterns []string +} + +func (rm *RecordingManager) RecordInteraction(snapshotID string, response *ResponseSnapshot) { + // Retrieve immutable request snapshot + request, exists := rm.storage.Get(snapshotID) + if !exists { + log.Printf("Request snapshot not found: %s", snapshotID) + return + } + + // Record with guaranteed data integrity + rm.recorder.RecordInteraction(request, response) +} + +func (r *Recorder) RecordInteraction(req *RequestSnapshot, res *ResponseSnapshot) error { + // Generate .http file with complete data + var buf bytes.Buffer + + // Write request + fmt.Fprintf(&buf, "### %s %s\n", req.Method, req.URL.String()) + fmt.Fprintf(&buf, "%s %s\n", req.Method, req.URL.String()) + fmt.Fprintf(&buf, "Host: %s\n", req.Host) + + for k, vv := range req.Headers { + for _, v := range vv { + fmt.Fprintf(&buf, "%s: %s\n", k, v) + } + } + + buf.WriteString("\n") + buf.Write(req.Body) + buf.WriteString("\n\n") + + // Write response + buf.WriteString("> {% \n") + fmt.Fprintf(&buf, " // Response: %d %s\n", res.StatusCode, http.StatusText(res.StatusCode)) + buf.WriteString(" // Headers:\n") + + for k, vv := range res.Headers { + for _, v := range vv { + fmt.Fprintf(&buf, " // %s: %s\n", k, v) + } + } + + buf.WriteString("%}\n\n") + + if len(res.Body) > 0 { + buf.WriteString("/*\n") + buf.Write(res.Body) + buf.WriteString("\n*/\n") + } else { + buf.WriteString("// [Binary response body: 0 bytes]\n") + } + + // Write to file + return r.writeToFile(buf.Bytes(), req, res) +} +``` + +## Migration Strategy + +### Phase 1: Introduce Snapshot System +- Add SnapshotMiddleware as first middleware +- Maintain existing recording system for compatibility +- Gradual migration of recording points + +### Phase 2: Update Mirror System +- Modify MirrorMiddleware to use snapshots +- Ensure mirror requests use snapshot data +- Test parity between old and new systems + +### Phase 3: Consolidate Recording +- Replace existing recording middleware +- Unified recording system using context-bound snapshots +- Remove duplicate body reading code + +### Phase 4: Cleanup +- Remove legacy recording code +- Optimize memory usage with sync.Pool +- Performance validation on target hardware (Pi Zero) + +## Benefits + +1. **Zero Extra Disk IO**: Protecs MicroSD by avoiding snapshot disk persistence +2. **Memory Efficiency**: Natural lifecycle tied to Request Context +3. **Data Integrity**: Request data is captured once and remains immutable +4. **Consistency**: All consumers use identical request data +5. **Traceability**: Clear lineage from original request to all recordings +6. **Performance**: Reduces duplicate body reads and re-cloning + +## Implementation Considerations + +### Memory Management +- Use `sync.Pool` for byte buffers +- Strict size limits on captured bodies +- Rely on GC for snapshot cleanup + +### Performance Impact +- Single body read vs multiple reads (net positive) +- Memory overhead for snapshot storage (manageable) +- Context propagation overhead (minimal) + +### Backward Compatibility +- Maintain existing .http file format +- Preserve existing API contracts +- Gradual migration path + +## Testing Strategy + +### Unit Tests +- Snapshot creation and immutability +- Response recording accuracy +- Memory cleanup verification + +### Integration Tests +- End-to-end request/response recording +- Mirror functionality with snapshots +- Parity validation between old/new systems + +### Performance Tests +- Memory usage comparison +- Throughput impact analysis +- Large request body handling + +## Future Enhancements + +1. **Compression**: Compress stored snapshots for memory efficiency +2. **Streaming**: Support for streaming request/response bodies +3. **Filtering**: Selective snapshot creation based on patterns +4. **Analytics**: Request/response analysis and metrics +5. **Export**: Snapshot export for debugging and analysis + +## Conclusion + +This snapshot-based approach provides a robust foundation for reliable request recording while solving the current issues with body consumption and data inconsistency. The phased implementation ensures minimal disruption while delivering immediate benefits. diff --git a/docs/SUMMARY.md b/docs/SUMMARY.md index 7061ab8..d750253 100644 --- a/docs/SUMMARY.md +++ b/docs/SUMMARY.md @@ -40,6 +40,7 @@ * [Feature Mapping](reference/FEATURE-MAPPING.md) ## Concepts +* [Request Recording](REQUEST_RECORDING_CONCEPT.md) * [Spotify Priming Strategy](concepts/spotify-priming-strategy.md) * [Spotify OAuth](concepts/spotify-oauth.md) diff --git a/pkg/discovery/dns_test.go b/pkg/discovery/dns_test.go index 82826ef..013d307 100644 --- a/pkg/discovery/dns_test.go +++ b/pkg/discovery/dns_test.go @@ -39,21 +39,8 @@ func TestDNSDiscovery_Interception(t *testing.T) { } // Test intercepting streamingoauth.bose.com - m3 := new(dns.Msg) - m3.SetQuestion("streamingoauth.bose.com.", dns.TypeA) - rw3 := &mockResponseWriter{} - d.ServeDNS(rw3, m3) - - if rw3.msg == nil || len(rw3.msg.Answer) == 0 { - t.Fatal("Expected response for streamingoauth.bose.com") - } - - if a, ok := rw3.msg.Answer[0].(*dns.A); ok { - if a.A.String() != serviceIP { - t.Errorf("Expected intercepted IP %s for streamingoauth.bose.com, got %s", serviceIP, a.A.String()) - } - } else { - t.Errorf("Expected A record for streamingoauth.bose.com, got %T", rw3.msg.Answer[0]) + if !d.shouldIntercept("streamingoauth.bose.com") { + t.Error("Expected streamingoauth.bose.com to be intercepted") } // Test aftertouch.test diff --git a/pkg/models/models.go b/pkg/models/models.go index c702928..8cf2de8 100644 --- a/pkg/models/models.go +++ b/pkg/models/models.go @@ -177,16 +177,25 @@ type ConfiguredSource struct { // ServiceDeviceInfo represents information about a SoundTouch device. type ServiceDeviceInfo struct { - DeviceID string `json:"device_id" xml:"deviceID,attr"` - ProductCode string `json:"product_code" xml:"type"` - DeviceSerialNumber string `json:"device_serial_number" xml:"serialnumber"` - ProductSerialNumber string `json:"product_serial_number" xml:"product_serial_number"` - FirmwareVersion string `json:"firmware_version" xml:"softwareVersion"` - IPAddress string `json:"ip_address" xml:"ipAddress"` - Name string `json:"name" xml:"name"` - MacAddress string `json:"mac_address,omitempty" xml:"-"` - DiscoveryMethod string `json:"discovery_method,omitempty"` - AccountID string `json:"account_id,omitempty"` + DeviceID string `json:"device_id" xml:"deviceID,attr"` + ProductCode string `json:"product_code" xml:"type"` + DeviceSerialNumber string `json:"device_serial_number" xml:"serialnumber"` + ProductSerialNumber string `json:"product_serial_number" xml:"product_serial_number"` + FirmwareVersion string `json:"firmware_version" xml:"softwareVersion"` + IPAddress string `json:"ip_address" xml:"ipAddress"` + Name string `json:"name" xml:"name"` + MacAddress string `json:"mac_address,omitempty" xml:"-"` + DiscoveryMethod string `json:"discovery_method,omitempty"` + AccountID string `json:"account_id,omitempty"` + Components []ServiceComponent `json:"components,omitempty" xml:"-"` +} + +// ServiceComponent represents a hardware or software component of a device. +type ServiceComponent struct { + Type string `xml:"type,attr"` + Category string `xml:"category,attr"` + SoftwareVersion string `xml:"softwareVersion"` + SerialNumber string `xml:"serialNumber"` } // CustomerSupportDevice represents device information for customer support purposes. diff --git a/pkg/service/datastore/datastore.go b/pkg/service/datastore/datastore.go index 89d0617..1352a43 100644 --- a/pkg/service/datastore/datastore.go +++ b/pkg/service/datastore/datastore.go @@ -290,12 +290,18 @@ func (ds *DataStore) parseDeviceInfoFile(path string) (*models.ServiceDeviceInfo deviceInfo := &models.ServiceDeviceInfo{ DeviceID: info.DeviceID, - ProductCode: fmt.Sprintf("%s %s", info.Type, info.ModuleType), + ProductCode: info.Type, Name: info.Name, DiscoveryMethod: info.DiscoveryMethod, } for _, comp := range info.Components { + deviceInfo.Components = append(deviceInfo.Components, models.ServiceComponent{ + Category: comp.Category, + SoftwareVersion: comp.SoftwareVersion, + SerialNumber: comp.SerialNumber, + }) + switch comp.Category { case "SCM": deviceInfo.FirmwareVersion = comp.SoftwareVersion @@ -453,9 +459,17 @@ func (ds *DataStore) GetRecents(account, device string) ([]models.ServiceRecent, } recents := []models.ServiceRecent{} + maxID := 0 for i := range recentsWrap.Recents { r := &recentsWrap.Recents[i] + + if id, err := strconv.Atoi(r.ID); err == nil { + if id > maxID { + maxID = id + } + } + recents = append(recents, models.ServiceRecent{ ServiceContentItem: models.ServiceContentItem{ ID: r.ID, @@ -472,6 +486,14 @@ func (ds *DataStore) GetRecents(account, device string) ([]models.ServiceRecent, }) } + // Ensure all recents have unique numeric IDs + for i := range recents { + if _, err := strconv.Atoi(recents[i].ID); err != nil || recents[i].ID == "" { + maxID++ + recents[i].ID = strconv.Itoa(maxID) + } + } + return recents, nil } @@ -867,6 +889,7 @@ type Settings struct { DNSBindAddr string `json:"dns_bind_addr,omitempty"` MirrorEnabled bool `json:"mirror_enabled"` MirrorEndpoints []string `json:"mirror_endpoints,omitempty"` + PreferredSource string `json:"preferred_source,omitempty"` InternalPaths []string `json:"internal_paths,omitempty"` Shortcuts map[string]int `json:"shortcuts,omitempty"` } diff --git a/pkg/service/handlers/handlers_marge.go b/pkg/service/handlers/handlers_marge.go index 1ff6a1b..8dbbe9e 100644 --- a/pkg/service/handlers/handlers_marge.go +++ b/pkg/service/handlers/handlers_marge.go @@ -61,7 +61,7 @@ func (s *Server) HandleMargePowerOn(w http.ResponseWriter, r *http.Request) { body, err := io.ReadAll(r.Body) if err != nil { log.Printf("[Marge] Failed to read power_on body: %v", err) - w.WriteHeader(http.StatusOK) // Silent failure is usually better for device requests + w.WriteHeader(http.StatusOK) return } @@ -191,15 +191,21 @@ func (s *Server) HandleMargeSoftwareUpdate(w http.ResponseWriter, r *http.Reques // This route is specifically used by firmware like Bose_Lisa/27.0.6. if chi.URLParam(r, "account") != "" { w.Header().Set("Content-Type", "application/vnd.bose.streaming-v1.2+xml") - _, _ = w.Write([]byte(marge.SoftwareUpdateToXML())) + + xmlData := marge.SoftwareUpdateToXML() + w.Header().Set("Content-Length", strconv.Itoa(len(xmlData))) + _, _ = w.Write([]byte(xmlData)) return } if len(swUpdateXML) > 0 { + w.Header().Set("Content-Length", strconv.Itoa(len(swUpdateXML))) _, _ = w.Write(swUpdateXML) } else { - _, _ = w.Write([]byte(marge.SoftwareUpdateToXML())) + xmlData := marge.SoftwareUpdateToXML() + w.Header().Set("Content-Length", strconv.Itoa(len(xmlData))) + _, _ = w.Write([]byte(xmlData)) } } @@ -415,7 +421,5 @@ func (s *Server) HandleMargeCustomerSupport(w http.ResponseWriter, r *http.Reque }, } s.ds.AddDeviceEvent(req.Device.ID, event) - - w.Header().Set("Content-Type", "application/vnd.bose.streaming-v1.2+xml") w.WriteHeader(http.StatusOK) } diff --git a/pkg/service/handlers/handlers_marge_test.go b/pkg/service/handlers/handlers_marge_test.go index 18784d8..015e2c0 100644 --- a/pkg/service/handlers/handlers_marge_test.go +++ b/pkg/service/handlers/handlers_marge_test.go @@ -54,10 +54,13 @@ func TestMargeSoftwareUpdate(t *testing.T) { } body, _ := io.ReadAll(res.Body) - // Should contain software_update or INDEX (if swupdate.xml exists) - if !strings.Contains(string(body), "software_update") && !strings.Contains(string(body), "INDEX") { + // Should contain INDEX as we updated swupdate.xml + if !strings.Contains(string(body), "INDEX") { t.Errorf("Unexpected response: %s", string(body)) } + if !strings.Contains(string(body), "0x0933") { + t.Errorf("Response missing VideoWave (0x0933) info: %s", string(body)) + } } func TestMargeAccountFull(t *testing.T) { @@ -761,6 +764,12 @@ func TestMargeAdvancedFeatures(t *testing.T) { if !strings.Contains(string(body), "123") { t.Errorf("Response body missing account ID: %s", body) } + if !strings.Contains(string(body), "ELIGIBLE_FOR_TRIAL") { + t.Errorf("Response body missing ELIGIBLE_FOR_TRIAL: %s", body) + } + if !strings.Contains(string(body), "STREAMING_QUALITY") { + t.Errorf("Response body missing STREAMING_QUALITY: %s", body) + } }) t.Run("StreamingToken", func(t *testing.T) { @@ -823,6 +832,10 @@ func TestMargeAdvancedFeatures(t *testing.T) { t.Errorf("Expected status OK, got %v", res.Status) } + if ct := res.Header.Get("Content-Type"); ct != "" { + t.Errorf("Expected no Content-Type for customer support upload (empty body), got %v", ct) + } + // Verify event was recorded events := ds.GetDeviceEvents("587A628A4042") found := false @@ -843,4 +856,40 @@ func TestMargeAdvancedFeatures(t *testing.T) { t.Error("Customer support event not found in event log") } }) + + t.Run("AddRecent_Reproduction", func(t *testing.T) { + account := "3230304" + device := "A81B6A536A98" + + // Setup sources for this device + deviceDir := ds.AccountDeviceDir(account, device) + _ = os.MkdirAll(deviceDir, 0755) + _ = os.WriteFile(filepath.Join(deviceDir, "Recents.xml"), []byte(""), 0644) + // No Sources.xml + + path := "/marge/streaming/account/" + account + "/device/" + device + "/recent" + payload := `2026-02-25T23:03:14+00:0010863533My top tracks playlist/playback/container/c3BvdGlmeTpwbGF5bGlzdDo3YklIMERKRUdoVjFSZ2duandOYWxntracklisturl` + + res, err := http.Post(ts.URL+path, "application/xml", strings.NewReader(payload)) + if err != nil { + t.Fatal(err) + } + defer res.Body.Close() + + if res.StatusCode != http.StatusCreated { + body, _ := io.ReadAll(res.Body) + t.Errorf("Expected status Created (201), got %v: %s", res.Status, body) + } + + // Verify it was saved + recents, err := ds.GetRecents(account, device) + if err != nil { + t.Fatalf("Failed to get recents: %v", err) + } + if len(recents) == 0 { + t.Error("Recents list is empty") + } else if recents[0].Name != "My top tracks playlist" { + t.Errorf("Expected name 'My top tracks playlist', got '%s'", recents[0].Name) + } + }) } diff --git a/pkg/service/handlers/handlers_setup.go b/pkg/service/handlers/handlers_setup.go index 7ccae3b..e77a6e1 100644 --- a/pkg/service/handlers/handlers_setup.go +++ b/pkg/service/handlers/handlers_setup.go @@ -156,6 +156,7 @@ func (s *Server) HandleGetSettings(w http.ResponseWriter, _ *http.Request) { dnsBindAddr := s.dnsBindAddr mirrorEnabled := s.mirrorEnabled mirrorEndpoints := s.mirrorEndpoints + preferredSource := s.preferredSource internalPaths := s.internalPaths enableSoundcorkProxy := s.enableSoundcorkProxy redact, logBody, record := s.proxyRedact, s.proxyLogBody, s.recordEnabled @@ -178,6 +179,7 @@ func (s *Server) HandleGetSettings(w http.ResponseWriter, _ *http.Request) { "dns_bind_addr": dnsBindAddr, "mirror_enabled": mirrorEnabled, "mirror_endpoints": mirrorEndpoints, + "preferred_source": preferredSource, "internal_paths": internalPaths, "enable_soundcork_proxy": enableSoundcorkProxy, "redact_logs": redact, @@ -203,6 +205,7 @@ func (s *Server) HandleUpdateSettings(w http.ResponseWriter, r *http.Request) { DNSBindAddr string `json:"dns_bind_addr"` MirrorEnabled bool `json:"mirror_enabled"` MirrorEndpoints []string `json:"mirror_endpoints"` + PreferredSource string `json:"preferred_source"` InternalPaths []string `json:"internal_paths"` EnableSoundcorkProxy bool `json:"enable_soundcork_proxy"` Shortcuts map[string]int `json:"shortcuts"` @@ -252,6 +255,7 @@ func (s *Server) HandleUpdateSettings(w http.ResponseWriter, r *http.Request) { s.mirrorEnabled = settings.MirrorEnabled s.mirrorEndpoints = settings.MirrorEndpoints + s.preferredSource = settings.PreferredSource s.internalPaths = settings.InternalPaths s.enableSoundcorkProxy = settings.EnableSoundcorkProxy @@ -285,6 +289,7 @@ func (s *Server) HandleUpdateSettings(w http.ResponseWriter, r *http.Request) { DNSBindAddr: s.dnsBindAddr, MirrorEnabled: s.mirrorEnabled, MirrorEndpoints: s.mirrorEndpoints, + PreferredSource: s.preferredSource, InternalPaths: s.internalPaths, EnableSoundcorkProxy: s.enableSoundcorkProxy, Shortcuts: s.shortcuts, diff --git a/pkg/service/handlers/mirror_middleware.go b/pkg/service/handlers/mirror_middleware.go index 3aa0e95..9795797 100644 --- a/pkg/service/handlers/mirror_middleware.go +++ b/pkg/service/handlers/mirror_middleware.go @@ -22,92 +22,153 @@ import ( // MirrorMiddleware returns a middleware that mirrors specific requests to the Bose upstream. func (s *Server) MirrorMiddleware(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - s.mu.RLock() - enabled := s.mirrorEnabled - endpoints := s.mirrorEndpoints - s.mu.RUnlock() + enabled, endpoints, preferredSource := s.getMirrorSettings() - if !enabled || len(endpoints) == 0 { + if !enabled || len(endpoints) == 0 || !s.shouldMirror(r.URL.Path, endpoints) { next.ServeHTTP(w, r) return } - shouldMirror := false + // Try to fetch snapshot from context + var snapshot *RequestSnapshot + if snap, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot); ok { + snapshot = snap + } - for _, pattern := range endpoints { - if matchPattern(pattern, r.URL.Path) { - shouldMirror = true - break + // Buffer request body if snapshot is missing (compatibility mode) + var bodyBytes []byte + if snapshot != nil { + bodyBytes = snapshot.Body + } else if r.Body != nil { + bodyBytes, _ = io.ReadAll(r.Body) + r.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + } + + // Use request context but detach it for background operations to prevent cancellation when the primary request finishes + detachedCtx := context.WithoutCancel(r.Context()) + if snapshot != nil { + detachedCtx = context.WithValue(detachedCtx, SnapshotKey, snapshot) + } + + if preferredSource == "upstream" { + s.mirrorUpstreamPreferred(detachedCtx, w, r, next, bodyBytes) + return + } + + s.mirrorLocalPreferred(detachedCtx, w, r, next, bodyBytes) + }) +} + +func (s *Server) getMirrorSettings() (bool, []string, string) { + s.mu.RLock() + defer s.mu.RUnlock() + + return s.mirrorEnabled, s.mirrorEndpoints, s.preferredSource +} + +func (s *Server) shouldMirror(path string, endpoints []string) bool { + for _, pattern := range endpoints { + if matchPattern(pattern, path) { + return true + } + } + + return false +} + +func (s *Server) mirrorUpstreamPreferred(detachedCtx context.Context, w http.ResponseWriter, r *http.Request, next http.Handler, bodyBytes []byte) { + log.Printf("[MIRROR] Upstream is preferred source for %s %s", r.Method, r.URL.Path) + + // Clone request for local execution + rLocal := r.Clone(detachedCtx) + rLocal.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + + localRecorder := &mirrorResponseRecorder{ + headers: make(http.Header), + body: &bytes.Buffer{}, + } + + // Run local handler in background + localDone := make(chan struct{}) + + go func() { + next.ServeHTTP(localRecorder, rLocal) + close(localDone) + }() + + // Clone request for mirror execution + rMirror := r.Clone(detachedCtx) + rMirror.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + + // Execute mirror synchronously + mirrorRes := s.performMirror(rMirror) + + // Send mirror response to client + if mirrorRes != nil && mirrorRes.status != 0 && mirrorRes.status < 500 { + for k, vv := range mirrorRes.headers { + for _, v := range vv { + w.Header().Add(k, v) } } - if !shouldMirror { - next.ServeHTTP(w, r) - return + w.WriteHeader(mirrorRes.status) + _, _ = w.Write(mirrorRes.body.Bytes()) + } else { + // Fallback to local if mirror failed + log.Printf("[MIRROR_ERR] Mirror failed, falling back to local for %s", r.URL.Path) + <-localDone + + for k, vv := range localRecorder.headers { + for _, v := range vv { + w.Header().Add(k, v) + } } - // Buffer request body for both local and mirror - var bodyBytes []byte - if r.Body != nil { - bodyBytes, _ = io.ReadAll(r.Body) - _ = r.Body.Close() + if localRecorder.status == 0 { + localRecorder.status = http.StatusOK } - // Prepare local request - r.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + w.WriteHeader(localRecorder.status) + _, _ = w.Write(localRecorder.body.Bytes()) + } - // Wrap response writer to capture local response for parity check - localRecorder := &mirrorResponseRecorder{ - headers: make(http.Header), - body: &bytes.Buffer{}, + // Perform parity check once local is done + go func() { + <-localDone + + if mirrorRes != nil { + s.checkParity(r, localRecorder, mirrorRes) } + }() +} - // Use a multi-writer if RecordMiddleware isn't already doing this, - // but let's just wrap it. +func (s *Server) mirrorLocalPreferred(detachedCtx context.Context, w http.ResponseWriter, r *http.Request, next http.Handler, bodyBytes []byte) { + // Default: local is preferred source of truth + // Prepare local request + r.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) - wrappedWriter := &parityResponseWriter{ - ResponseWriter: w, - recorder: localRecorder, - } + // Wrap response writer to capture local response for parity check + localRecorder := &mirrorResponseRecorder{ + headers: make(http.Header), + body: &bytes.Buffer{}, + } - if r.Method == http.MethodGet { - // GET: Local is primary, Mirror is asynchronous - log.Printf("[MIRROR] Mirroring GET %s asynchronously", r.URL.Path) + wrappedWriter := &parityResponseWriter{ + ResponseWriter: w, + recorder: localRecorder, + } - // We need a clone for the async call, detached from original request context - // We use context.Background() because the original request's context - // will be canceled as soon as the local handler finishes and returns - // the response to the speaker. - //nolint:contextcheck - rMirror := r.Clone(context.Background()) - rMirror.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + log.Printf("[MIRROR] Mirroring %s %s %s", r.Method, r.URL.Path, map[bool]string{true: "asynchronously", false: "synchronously"}[r.Method == http.MethodGet]) - // For GET, we run mirror in background and don't wait for parity in real-time - // or we can wait for local to finish then trigger parity asynchronously. + rMirror := r.Clone(detachedCtx) + rMirror.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) - next.ServeHTTP(wrappedWriter, r) + next.ServeHTTP(wrappedWriter, r) - go func() { - mirrorRes := s.performMirror(rMirror) - s.checkParity(r, localRecorder, mirrorRes) - }() - } else { - // POST/PUT/DELETE: Local is primary for speaker response, but we sync synchronously - log.Printf("[MIRROR] Mirroring %s %s synchronously", r.Method, r.URL.Path) - - // We need a clone for the background sync call - //nolint:contextcheck - rMirror := r.Clone(context.Background()) - rMirror.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) - - next.ServeHTTP(wrappedWriter, r) - - go func() { - mirrorRes := s.performMirror(rMirror) - s.checkParity(r, localRecorder, mirrorRes) - }() - } - }) + go func() { + mirrorRes := s.performMirror(rMirror) + s.checkParity(r, localRecorder, mirrorRes) + }() } type parityResponseWriter struct { @@ -142,6 +203,41 @@ func (p *parityResponseWriter) WriteHeader(statusCode int) { } func (s *Server) performMirror(r *http.Request) *mirrorResponseRecorder { + // Try to fetch snapshot from context + var snapshot *RequestSnapshot + if snap, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot); ok { + snapshot = snap + } + + // Preserve request body for recording before it gets consumed by the proxy + var requestForRecording *http.Request + if s.recorder != nil && s.recordEnabled { + requestForRecording = r.Clone(r.Context()) + if snapshot != nil { + // Use snapshot for both proxy and recording + r.Body = io.NopCloser(bytes.NewReader(snapshot.Body)) + requestForRecording.Body = io.NopCloser(bytes.NewReader(snapshot.Body)) + } else if r.Body != nil { + // Compatibility fallback + bodyBytes, err := io.ReadAll(r.Body) + if err != nil { + log.Printf("[MIRROR_ERR] Failed to read request body for recording: %v", err) + } else { + // Restore body for proxy + r.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + // Set body for recording + requestForRecording.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + } + } + + // Ensure Content-Length is set for the recording clone + if requestForRecording.Body != nil { + if snapshot != nil { + requestForRecording.ContentLength = int64(len(snapshot.Body)) + } + } + } + host := r.Host if host == "" || host == "localhost" { host = "streaming.bose.com" @@ -183,9 +279,9 @@ func (s *Server) performMirror(r *http.Request) *mirrorResponseRecorder { proxy.ModifyResponse = func(res *http.Response) error { res.Header.Set("X-Proxy-Origin", "upstream-mirror") - // Record mirrored interaction - if s.recorder != nil && s.recordEnabled { - _ = s.recorder.Record("mirror", r, res) + // Record mirrored interaction with preserved request body + if s.recorder != nil && s.recordEnabled && requestForRecording != nil { + _ = s.recorder.Record("mirror", requestForRecording, res) } return nil diff --git a/pkg/service/handlers/mirror_preferred_test.go b/pkg/service/handlers/mirror_preferred_test.go new file mode 100644 index 0000000..c53000e --- /dev/null +++ b/pkg/service/handlers/mirror_preferred_test.go @@ -0,0 +1,160 @@ +package handlers + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "strings" + "testing" + + "github.com/gesellix/bose-soundtouch/pkg/service/datastore" +) + +func TestMirrorMiddleware_PreferredSource(t *testing.T) { + tempDir, err := os.MkdirTemp("", "mirror-test") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer os.RemoveAll(tempDir) + + ds := datastore.NewDataStore(tempDir) + _ = ds.Initialize() + + // 1. Setup local handler + r := http.NewServeMux() + r.HandleFunc("/test/local", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Source", "local") + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("local response")) + }) + + // 2. Setup "upstream" mock server + upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("X-Source", "upstream") + w.WriteHeader(http.StatusCreated) + _, _ = w.Write([]byte("upstream response")) + })) + defer upstreamServer.Close() + + // 3. Setup our server with MirrorMiddleware + server := NewServer(ds, nil, "http://localhost:8000", false, false, false, false, false, false) + server.SetMirrorSettings(true, []string{"/test/local"}, "local") + + // We need to trick performMirror to use our mock upstream. + // performMirror uses r.Host. + upstreamURL := upstreamServer.URL + upstreamHost := strings.TrimPrefix(upstreamURL, "http://") + + middleware := server.MirrorMiddleware(r) + + t.Run("PreferredLocal", func(t *testing.T) { + server.SetMirrorSettings(true, []string{"/test/local"}, "local") + + req := httptest.NewRequest("GET", "/test/local", nil) + req.Host = upstreamHost // So performMirror targets the mock upstream + w := httptest.NewRecorder() + + middleware.ServeHTTP(w, req) + + if w.Code != http.StatusOK { + t.Errorf("Expected status 200, got %d", w.Code) + } + if w.Header().Get("X-Source") != "local" { + t.Errorf("Expected X-Source: local, got %s", w.Header().Get("X-Source")) + } + if w.Body.String() != "local response" { + t.Errorf("Expected 'local response', got '%s'", w.Body.String()) + } + }) + + t.Run("PreferredUpstream", func(t *testing.T) { + server.SetMirrorSettings(true, []string{"/test/local"}, "upstream") + + req := httptest.NewRequest("GET", "/test/local", nil) + req.Host = upstreamHost + w := httptest.NewRecorder() + + middleware.ServeHTTP(w, req) + + if w.Code != http.StatusCreated { + t.Errorf("Expected status 201, got %d", w.Code) + } + if w.Header().Get("X-Source") != "upstream" { + t.Errorf("Expected X-Source: upstream, got %s", w.Header().Get("X-Source")) + } + if w.Body.String() != "upstream response" { + t.Errorf("Expected 'upstream response', got '%s'", w.Body.String()) + } + }) + + t.Run("FallbackToLocal", func(t *testing.T) { + server.SetMirrorSettings(true, []string{"/test/local"}, "upstream") + + // Use a non-existent host for mirror to trigger failure + req := httptest.NewRequest("GET", "/test/local", nil) + req.Host = "nonexistent.invalid" + w := httptest.NewRecorder() + + middleware.ServeHTTP(w, req) + + // Should fallback to local + if w.Code != http.StatusOK { + t.Errorf("Expected status 200 (fallback), got %d", w.Code) + } + if w.Header().Get("X-Source") != "local" { + t.Errorf("Expected X-Source: local (fallback), got %s", w.Header().Get("X-Source")) + } + }) +} + +func TestSettingsAPI_PreferredSource(t *testing.T) { + tempDir, err := os.MkdirTemp("", "settings-test") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer os.RemoveAll(tempDir) + + ds := datastore.NewDataStore(tempDir) + _ = ds.Initialize() + + server := NewServer(ds, nil, "http://localhost:8000", false, false, false, false, false, false) + + // Test GET initial + req := httptest.NewRequest("GET", "/setup/settings", nil) + w := httptest.NewRecorder() + server.HandleGetSettings(w, req) + + var settings map[string]interface{} + json.Unmarshal(w.Body.Bytes(), &settings) + if settings["preferred_source"] != "" && settings["preferred_source"] != "local" { + t.Errorf("Initial preferred_source unexpected: %v", settings["preferred_source"]) + } + + // Test UPDATE + update := map[string]interface{}{ + "preferred_source": "upstream", + } + body, err := json.Marshal(update) + if err != nil { + t.Fatalf("Failed to marshal update: %v", err) + } + req = httptest.NewRequest("POST", "/setup/settings", bytes.NewBuffer(body)) + w = httptest.NewRecorder() + server.HandleUpdateSettings(w, req) + + if w.Code != http.StatusOK { + t.Errorf("POST /setup/settings failed: %d", w.Code) + } + + if server.preferredSource != "upstream" { + t.Errorf("Server preferredSource did not update: %s", server.preferredSource) + } + + // Verify persistence + persisted, _ := ds.GetSettings() + if persisted.PreferredSource != "upstream" { + t.Errorf("Datastore did not persist PreferredSource: %s", persisted.PreferredSource) + } +} diff --git a/pkg/service/handlers/mirror_test.go b/pkg/service/handlers/mirror_test.go index f179df0..8a33d53 100644 --- a/pkg/service/handlers/mirror_test.go +++ b/pkg/service/handlers/mirror_test.go @@ -2,6 +2,8 @@ package handlers import ( "encoding/json" + "fmt" + "io" "net/http" "net/http/httptest" "os" @@ -43,7 +45,7 @@ func TestMirroring(t *testing.T) { recorder := proxy.NewRecorder(tempDir) server.SetRecorder(recorder) server.SetRecordEnabled(true) - server.SetMirrorSettings(true, []string{"/streaming/account/*/device/*/recent"}) + server.SetMirrorSettings(true, []string{"/streaming/account/*/device/*/recent"}, "local") ts := httptest.NewServer(r) defer ts.Close() @@ -137,6 +139,82 @@ func TestMirroring(t *testing.T) { t.Errorf("Expected Upstream Content-Type application/vnd.bose.streaming-v1.2+xml, got %s", ct) } }) + + t.Run("POST Request Body Preservation", func(t *testing.T) { + // Set recorder to synchronous mode for testing + os.Setenv("RECORDER_ASYNC", "false") + defer os.Unsetenv("RECORDER_ASYNC") + + // Create a mock upstream that echoes back the request body + postUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == "POST" && strings.HasSuffix(r.URL.Path, "/scmudc/A81B6A536A98") { + // Read the request body + body, err := io.ReadAll(r.Body) + if err != nil { + w.WriteHeader(http.StatusInternalServerError) + return + } + + // Echo back the body in response for verification + w.Header().Set("Content-Type", "application/json") + w.Header().Set("X-Request-Body-Length", fmt.Sprintf("%d", len(body))) + w.WriteHeader(http.StatusOK) + _, _ = w.Write(body) + return + } + w.WriteHeader(http.StatusNotFound) + })) + defer postUpstream.Close() + + // Setup mirroring for the POST endpoint + server.SetMirrorSettings(true, []string{"/v1/scmudc/*"}, "local") + + requestBody := `{"envelope":{"monoTime":234906,"payloadProtocolVersion":"3.1","payloadType":"scmudc","protocolVersion":"1.0","time":"2026-02-25T23:03:14.976349+00:00","uniqueId":"A81B6A536A98"},"payload":{"deviceInfo":{"boseID":"3230304","deviceID":"A81B6A536A98","deviceType":"SoundTouch 10","serialNumber":"I6332527703739342000020","softwareVersion":"27.0.6.46330.5043500 epdbuild.trunk.hepdswbld04.2022-08-04T11:20:29","systemSerialNumber":"069231P63364828AE"},"events":[{"data":{"play-state":"PAUSE_STATE"},"monoTime":234904,"time":"2026-02-25T23:03:14.973466+00:00","type":"play-state-changed"}]}}` + + path := "/v1/scmudc/A81B6A536A98" + req, _ := http.NewRequest("POST", ts.URL+path, strings.NewReader(requestBody)) + req.Header.Set("Content-Type", "text/json; charset=utf-8") + req.Host = strings.TrimPrefix(postUpstream.URL, "http://") + + res, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer res.Body.Close() + + if res.StatusCode != http.StatusOK { + t.Errorf("Expected status OK, got %v", res.Status) + } + + // Wait briefly for the synchronous recording to complete + time.Sleep(100 * time.Millisecond) + + // Check if the mirrored interaction was recorded with the request body + matchesMirror, _ := filepath.Glob(filepath.Join(tempDir, "interactions", "*", "mirror", "v1", "scmudc", "*", "*-POST.http")) + if len(matchesMirror) == 0 { + // Try broader search pattern + allHttpFiles, _ := filepath.Glob(filepath.Join(tempDir, "interactions", "*", "*", "*", "*", "*", "*.http")) + t.Errorf("Expected to find mirrored POST interaction. All .http files found: %v", allHttpFiles) + } else { + // Read the recorded mirrored interaction + recordedContent, err := os.ReadFile(matchesMirror[0]) + if err != nil { + t.Fatalf("Failed to read recorded mirror interaction: %v", err) + } + + recordedStr := string(recordedContent) + + // Check if the request body was preserved in the recording + if !strings.Contains(recordedStr, requestBody) { + t.Errorf("Request body not found in mirrored recording. Content: %s", recordedStr) + } + + // Check if the Content-Type header was preserved + if !strings.Contains(recordedStr, "Content-Type: text/json; charset=utf-8") { + t.Errorf("Content-Type header not found in mirrored recording. Content: %s", recordedStr) + } + } + }) } // SetRecordEnabled is a helper for testing diff --git a/pkg/service/handlers/recorder_middleware.go b/pkg/service/handlers/recorder_middleware.go index 7706238..f90747d 100644 --- a/pkg/service/handlers/recorder_middleware.go +++ b/pkg/service/handlers/recorder_middleware.go @@ -28,16 +28,18 @@ func (s *Server) RecordMiddleware(next http.Handler) http.Handler { } } - // Buffer the request body if it exists + // Use snapshot if available, otherwise buffer body (compatibility mode) + var snapshot *RequestSnapshot + if s, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot); ok { + snapshot = s + } + var reqBody []byte - - if r.Body != nil { - var err error - - reqBody, err = io.ReadAll(r.Body) - if err == nil { - r.Body = io.NopCloser(bytes.NewBuffer(reqBody)) - } + if snapshot != nil { + reqBody = snapshot.Body + } else if r.Body != nil { + reqBody, _ = io.ReadAll(r.Body) + r.Body = io.NopCloser(bytes.NewBuffer(reqBody)) } // wrap ResponseWriter to capture the response @@ -54,7 +56,7 @@ func (s *Server) RecordMiddleware(next http.Handler) http.Handler { defer func() { _ = res.Body.Close() }() } - // Put back the original request body for recording + // Restore body for recording r.Body = io.NopCloser(bytes.NewBuffer(reqBody)) _ = s.recorder.Record("self", r, res) diff --git a/pkg/service/handlers/server.go b/pkg/service/handlers/server.go index 858ec99..6a2c2bc 100644 --- a/pkg/service/handlers/server.go +++ b/pkg/service/handlers/server.go @@ -1,6 +1,7 @@ package handlers import ( + "bytes" "context" "fmt" "io" @@ -42,6 +43,7 @@ type Server struct { dnsBindAddr string mirrorEnabled bool mirrorEndpoints []string + preferredSource string internalPaths []string enableSoundcorkProxy bool shortcuts map[string]int @@ -59,6 +61,27 @@ type Server struct { spotifyService *spotify.Service } +// RequestSnapshot represents an immutable snapshot of an HTTP request. +type RequestSnapshot struct { + Method string + URL *url.URL + Headers http.Header + Body []byte + Host string + Timestamp time.Time +} + +type ctxKey struct{ name string } + +// SnapshotKey is the context key for the RequestSnapshot. +var SnapshotKey = &ctxKey{"request_snapshot"} + +var bufferPool = sync.Pool{ + New: func() interface{} { + return new(bytes.Buffer) + }, +} + // NewServer creates a new SoundTouch service server. func NewServer(ds *datastore.DataStore, sm *setup.Manager, serverURL string, proxyRedact, proxyLogBody, recordEnabled, enableSoundcorkProxy, migrationEnabled, migrationDryRun bool) *Server { // Initialize migration manager @@ -317,12 +340,13 @@ func (s *Server) SetMgmtConfig(username, password string) { } // SetMirrorSettings sets the mirroring settings for the server. -func (s *Server) SetMirrorSettings(enabled bool, endpoints []string) { +func (s *Server) SetMirrorSettings(enabled bool, endpoints []string, preferredSource string) { s.mu.Lock() defer s.mu.Unlock() s.mirrorEnabled = enabled s.mirrorEndpoints = endpoints + s.preferredSource = preferredSource } // SetInternalPaths sets the internal paths for the server. diff --git a/pkg/service/handlers/snapshot_integrity_test.go b/pkg/service/handlers/snapshot_integrity_test.go new file mode 100644 index 0000000..16efa4f --- /dev/null +++ b/pkg/service/handlers/snapshot_integrity_test.go @@ -0,0 +1,126 @@ +package handlers + +import ( + "bytes" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + "time" + + "github.com/gesellix/bose-soundtouch/pkg/service/datastore" + "github.com/gesellix/bose-soundtouch/pkg/service/proxy" +) + +func TestSnapshotIntegrity_SelfAndMirror(t *testing.T) { + tempDir, err := os.MkdirTemp("", "recording-test-*") + if err != nil { + t.Fatal(err) + } + defer os.RemoveAll(tempDir) + + ds := datastore.NewDataStore(tempDir) + recorder := proxy.NewRecorder(tempDir) + s := NewServer(ds, nil, "http://localhost:8000", false, false, true, false, false, false) + s.SetRecorder(recorder) + s.SetMirrorSettings(true, []string{"/mirror/*"}, "local") + + // Upstream mock + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + w.Header().Set("X-Request-Body-Length", fmt.Sprintf("%d", len(body))) + w.WriteHeader(http.StatusOK) + w.Write([]byte("upstream response")) + })) + defer upstream.Close() + + // Configure mirror to point to our mock upstream + s.SetMirrorSettings(true, []string{"/mirror/*"}, "local") + // We need to override the host in performMirror but for tests we can just mock it via env if needed or rely on the fact that performMirror uses r.Host + + handler := s.SnapshotMiddleware(s.MirrorMiddleware(s.RecordMiddleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + w.WriteHeader(http.StatusOK) + w.Write([]byte("local response: " + string(body))) + })))) + + bodyText := `{"test":"integrity"}` + req := httptest.NewRequest("POST", "http://localhost:8000/mirror/test", strings.NewReader(bodyText)) + req.Header.Set("Content-Type", "application/json") + // Override r.Host to point to our mock upstream (performMirror will use it) + req.Host = strings.TrimPrefix(upstream.URL, "http://") + + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + // Wait for async operations + time.Sleep(200 * time.Millisecond) + + var selfFile, mirrorFile string + _ = filepath.Walk(tempDir, func(path string, info os.FileInfo, err error) error { + if err != nil { + return err + } + if !info.IsDir() && strings.HasSuffix(path, ".http") { + if strings.Contains(path, "/self/") { + selfFile = path + } else if strings.Contains(path, "/mirror/") { + mirrorFile = path + } + } + return nil + }) + + // Retry a few times for async operations + for i := 0; i < 10 && (selfFile == "" || mirrorFile == ""); i++ { + time.Sleep(100 * time.Millisecond) + _ = filepath.Walk(tempDir, func(path string, info os.FileInfo, err error) error { + if err != nil { + return err + } + if !info.IsDir() && strings.HasSuffix(path, ".http") { + if strings.Contains(path, "/self/") { + selfFile = path + } else if strings.Contains(path, "/mirror/") { + mirrorFile = path + } + } + return nil + }) + } + + if selfFile == "" { + // Try one more scan + filepath.Walk(tempDir, func(path string, info os.FileInfo, err error) error { + if !info.IsDir() && strings.HasSuffix(path, ".http") { + if strings.Contains(path, "/self/") { + selfFile = path + } else if strings.Contains(path, "/mirror/") { + mirrorFile = path + } + } + return nil + }) + } + + if selfFile == "" { + t.Fatal("Self recording file not found") + } + if mirrorFile == "" { + t.Fatal("Mirror recording file not found") + } + + selfContent, _ := os.ReadFile(selfFile) + mirrorContent, _ := os.ReadFile(mirrorFile) + + if !bytes.Contains(selfContent, []byte(bodyText)) { + t.Errorf("Self recording missing body. Content:\n%s", string(selfContent)) + } + if !bytes.Contains(mirrorContent, []byte(bodyText)) { + t.Errorf("Mirror recording missing body. Content:\n%s", string(mirrorContent)) + } +} diff --git a/pkg/service/handlers/snapshot_middleware.go b/pkg/service/handlers/snapshot_middleware.go new file mode 100644 index 0000000..6b7d1e7 --- /dev/null +++ b/pkg/service/handlers/snapshot_middleware.go @@ -0,0 +1,89 @@ +package handlers + +import ( + "bytes" + "context" + "errors" + "io" + "net/http" + "net/url" + "time" +) + +// SnapshotMiddleware creates an immutable snapshot of the request body and metadata. +func (s *Server) SnapshotMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // 1. Check if we already have a snapshot (shouldn't happen with correct middleware order) + if _, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot); ok { + next.ServeHTTP(w, r) + return + } + + // 2. Capture body with size limit (e.g. 2MB) + const maxBodySize = 2 * 1024 * 1024 + + var body []byte + + if r.Body != nil { + buf, ok := bufferPool.Get().(*bytes.Buffer) + if !ok { + buf = new(bytes.Buffer) + } + + buf.Reset() + defer bufferPool.Put(buf) + + // Read up to maxBodySize + 1 to detect truncation + _, err := io.CopyN(buf, r.Body, maxBodySize+1) + _ = r.Body.Close() + + if err != nil && !errors.Is(err, io.EOF) { + // If reading fails, proceed with empty body but log it? + // For now, we follow the concept and proceed. + body = []byte{} + } else { + body = buf.Bytes() + if int64(len(body)) > maxBodySize { + body = body[:maxBodySize] + // Optional: mark as truncated if we add that field later + } + // Copy to a fresh byte slice because buf.Bytes() is a slice into the buffer + body = append([]byte(nil), body...) + } + } + + // 3. Create snapshot + snapshot := &RequestSnapshot{ + Method: r.Method, + URL: cloneURL(r.URL), + Headers: r.Header.Clone(), + Body: body, + Host: r.Host, + Timestamp: time.Now(), + } + + // 4. Inject into context + ctx := context.WithValue(r.Context(), SnapshotKey, snapshot) + r = r.WithContext(ctx) + + // 5. Restore r.Body for downstream compatibility + r.Body = io.NopCloser(bytes.NewReader(snapshot.Body)) + + next.ServeHTTP(w, r) + }) +} + +// cloneURL provides a deep copy of a URL. +func cloneURL(u *url.URL) *url.URL { + if u == nil { + return nil + } + + u2 := *u + if u.User != nil { + u2.User = new(url.Userinfo) + *u2.User = *u.User + } + + return &u2 +} diff --git a/pkg/service/handlers/snapshot_middleware_test.go b/pkg/service/handlers/snapshot_middleware_test.go new file mode 100644 index 0000000..bba6fe5 --- /dev/null +++ b/pkg/service/handlers/snapshot_middleware_test.go @@ -0,0 +1,101 @@ +package handlers + +import ( + "bytes" + "io" + "net/http" + "net/http/httptest" + "testing" +) + +func TestSnapshotMiddleware(t *testing.T) { + s := &Server{} + + t.Run("CapturesBodyAndMetadata", func(t *testing.T) { + bodyText := "hello world" + req := httptest.NewRequest("POST", "http://example.com/foo?bar=baz", bytes.NewBufferString(bodyText)) + req.Header.Set("Content-Type", "text/plain") + req.Host = "example.com" + + recorded := false + handler := s.SnapshotMiddleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + recorded = true + + // Verify snapshot in context + snapshot, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot) + if !ok { + t.Fatal("Snapshot not found in context") + } + + if snapshot.Method != "POST" { + t.Errorf("Expected method POST, got %s", snapshot.Method) + } + if snapshot.URL.Path != "/foo" { + t.Errorf("Expected path /foo, got %s", snapshot.URL.Path) + } + if snapshot.Headers.Get("Content-Type") != "text/plain" { + t.Errorf("Expected header text/plain, got %s", snapshot.Headers.Get("Content-Type")) + } + if string(snapshot.Body) != bodyText { + t.Errorf("Expected body %s, got %s", bodyText, string(snapshot.Body)) + } + if snapshot.Host != "example.com" { + t.Errorf("Expected host example.com, got %s", snapshot.Host) + } + + // Verify r.Body is still readable + body, _ := io.ReadAll(r.Body) + if string(body) != bodyText { + t.Errorf("Expected r.Body to be %s, got %s", bodyText, string(body)) + } + })) + + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + + if !recorded { + t.Error("Handler was not called") + } + }) + + t.Run("HandlesEmptyBody", func(t *testing.T) { + req := httptest.NewRequest("GET", "http://example.com/foo", nil) + + handler := s.SnapshotMiddleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + snapshot, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot) + if !ok { + t.Fatal("Snapshot not found in context") + } + if len(snapshot.Body) != 0 { + t.Errorf("Expected empty body, got %d bytes", len(snapshot.Body)) + } + })) + + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + }) + + t.Run("RespectsSizeLimit", func(t *testing.T) { + largeBody := make([]byte, 3*1024*1024) // 3MB + for i := range largeBody { + largeBody[i] = 'A' + } + + req := httptest.NewRequest("POST", "http://example.com/foo", bytes.NewReader(largeBody)) + + handler := s.SnapshotMiddleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + snapshot, ok := r.Context().Value(SnapshotKey).(*RequestSnapshot) + if !ok { + t.Fatal("Snapshot not found in context") + } + + const maxBodySize = 2 * 1024 * 1024 + if len(snapshot.Body) != maxBodySize { + t.Errorf("Expected body size %d, got %d", maxBodySize, len(snapshot.Body)) + } + })) + + rr := httptest.NewRecorder() + handler.ServeHTTP(rr, req) + }) +} diff --git a/pkg/service/handlers/static/swupdate.xml b/pkg/service/handlers/static/swupdate.xml index 151a14d..ffff0d9 100644 --- a/pkg/service/handlers/static/swupdate.xml +++ b/pkg/service/handlers/static/swupdate.xml @@ -4,46 +4,50 @@ - - - + + + - - - + + + - - - + + + - - - - - - - + + + + + + + - - + + - + @@ -51,216 +55,248 @@ - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + - - - + + + @@ -268,8 +304,8 @@ - - + + @@ -277,8 +313,8 @@ - - + + @@ -286,8 +322,9 @@ - - + + @@ -295,8 +332,9 @@ - - + + @@ -304,8 +342,8 @@ - - + + diff --git a/pkg/service/handlers/web/index.html b/pkg/service/handlers/web/index.html index 983559c..2947d4a 100644 --- a/pkg/service/handlers/web/index.html +++ b/pkg/service/handlers/web/index.html @@ -127,6 +127,9 @@ Enable Background Mirroring to Bose Cloud
+
diff --git a/pkg/service/handlers/web/js/script.js b/pkg/service/handlers/web/js/script.js index 4e7aa2b..85b9aca 100644 --- a/pkg/service/handlers/web/js/script.js +++ b/pkg/service/handlers/web/js/script.js @@ -155,6 +155,9 @@ async function fetchSettings() { if (settings.mirror_enabled !== undefined) { document.getElementById('mirror-enabled').checked = settings.mirror_enabled; } + if (settings.preferred_source !== undefined) { + document.getElementById('preferred-source-upstream').checked = settings.preferred_source === 'upstream'; + } if (settings.mirror_endpoints) { document.getElementById('mirror-endpoints').value = settings.mirror_endpoints.join('\n'); } @@ -226,6 +229,7 @@ async function updateSettings() { dns_upstream: document.getElementById('dns-upstream').value, dns_bind_addr: document.getElementById('dns-bind').value, mirror_enabled: document.getElementById('mirror-enabled').checked, + preferred_source: document.getElementById('preferred-source-upstream').checked ? 'upstream' : 'local', mirror_endpoints: document.getElementById('mirror-endpoints').value.split('\n').map(s => s.trim()).filter(s => s !== ''), internal_paths: document.getElementById('internal-paths').value.split('\n').map(s => s.trim()).filter(s => s !== ''), enable_soundcork_proxy: document.getElementById('enable-soundcork-proxy').checked diff --git a/pkg/service/marge/marge.go b/pkg/service/marge/marge.go index 034cb65..1602f29 100644 --- a/pkg/service/marge/marge.go +++ b/pkg/service/marge/marge.go @@ -3,10 +3,12 @@ package marge import ( + "bytes" "encoding/xml" "fmt" "os" "strconv" + "strings" "time" "github.com/gesellix/bose-soundtouch/pkg/models" @@ -66,16 +68,22 @@ func ConfiguredSourceToXML(cs models.ConfiguredSource) ([]byte, error) { Name string `xml:"name"` SourceProviderID string `xml:"sourceproviderid"` SourceName string `xml:"sourcename"` - SourceSettings string `xml:"sourcesettings"` + SourceSettings string `xml:"sourceSettings"` UpdatedOn string `xml:"updatedOn"` Username string `xml:"username"` } providerID := 0 + tokenType := "token" for i, p := range constants.Providers { if p == cs.SourceKeyType { providerID = i + 1 + + if p == "SPOTIFY" { + tokenType = "token_version_3" + } + break } } @@ -90,25 +98,41 @@ func ConfiguredSourceToXML(cs models.ConfiguredSource) ([]byte, error) { UpdatedOn: DateStr, Username: cs.SourceKeyAccount, } - sxml.Credential.Type = "token" + sxml.Credential.Type = tokenType sxml.Credential.Value = cs.Secret return xml.Marshal(sxml) } +// EscapeXML escapes special characters for XML. +func EscapeXML(s string) string { + var b bytes.Buffer + if err := xml.EscapeText(&b, []byte(s)); err != nil { + return s + } + + return b.String() +} + // GetConfiguredSourceXML returns the XML representation of a configured source as a string. func GetConfiguredSourceXML(cs models.ConfiguredSource) string { providerID := 0 + tokenType := "token" for i, p := range constants.Providers { if p == cs.SourceKeyType { providerID = i + 1 + + if p == "SPOTIFY" { + tokenType = "token_version_3" + } + break } } - return fmt.Sprintf(`%s%s%s%d%s%s%s`, - cs.ID, DateStr, cs.Secret, cs.SourceKeyAccount, providerID, cs.DisplayName, DateStr, cs.SourceKeyAccount) + return fmt.Sprintf(`%s%s%s%d%s%s%s`, + EscapeXML(cs.ID), DateStr, EscapeXML(tokenType), EscapeXML(cs.Secret), EscapeXML(cs.SourceKeyAccount), providerID, EscapeXML(cs.DisplayName), DateStr, EscapeXML(cs.SourceKeyAccount)) } // PresetsToXML converts account presets to XML format for Marge responses. @@ -127,12 +151,12 @@ func PresetsToXML(ds *datastore.DataStore, account, device string) ([]byte, erro for i := range presets { p := &presets[i] - res += fmt.Sprintf(``, p.ID) - res += fmt.Sprintf(`%s`, p.ContainerArt) - res += fmt.Sprintf(`%s`, p.Type) + res += fmt.Sprintf(``, EscapeXML(p.ID)) + res += fmt.Sprintf(`%s`, EscapeXML(p.ContainerArt)) + res += fmt.Sprintf(`%s`, EscapeXML(p.Type)) res += fmt.Sprintf(`%s`, DateStr) - res += fmt.Sprintf(`%s`, p.Location) - res += fmt.Sprintf(`%s`, p.Name) + res += fmt.Sprintf(`%s`, EscapeXML(p.Location)) + res += fmt.Sprintf(`%s`, EscapeXML(p.Name)) // Content Item Source for j := range sources { @@ -174,22 +198,30 @@ func RecentsToXML(ds *datastore.DataStore, account, device string) ([]byte, erro lastPlayed = time.Unix(sec, 0).Format(time.RFC3339) } - res += fmt.Sprintf(``, r.ID) - res += fmt.Sprintf(`%s`, r.Type) + res += fmt.Sprintf(``, EscapeXML(r.ID)) + res += fmt.Sprintf(`%s`, EscapeXML(r.Type)) res += fmt.Sprintf(`%s`, DateStr) - res += fmt.Sprintf(`%s`, lastPlayed) - res += fmt.Sprintf(`%s`, r.Location) - res += fmt.Sprintf(`%s`, r.Name) + res += fmt.Sprintf(`%s`, EscapeXML(lastPlayed)) + res += fmt.Sprintf(`%s`, EscapeXML(r.Location)) + res += fmt.Sprintf(`%s`, EscapeXML(r.Name)) // Content Item Source + sourceID := "" + for j := range sources { s := sources[j] if s.ID == r.SourceID || (s.SourceKeyType == r.Source && s.SourceKeyAccount == r.SourceAccount) { res += GetConfiguredSourceXML(s) + sourceID = s.ID + break } } + if sourceID != "" { + res += fmt.Sprintf(`%s`, EscapeXML(sourceID)) + } + res += fmt.Sprintf(`%s`, DateStr) res += `` } @@ -201,7 +233,20 @@ func RecentsToXML(ds *datastore.DataStore, account, device string) ([]byte, erro // ProviderSettingsToXML generates provider settings XML for the specified account. func ProviderSettingsToXML(account string) string { - return fmt.Sprintf(`%sELIGIBLE_FOR_TRIALtrue14`, account) + return xml.Header + fmt.Sprintf(` + + %s + ELIGIBLE_FOR_TRIAL + false + 14 + + + %s + STREAMING_QUALITY + 2 + 15 + + `, EscapeXML(account), EscapeXML(account)) } // SoftwareUpdateToXML generates software update configuration XML. @@ -218,7 +263,7 @@ func AccountFullToXML(ds *datastore.DataStore, account string) ([]byte, error) { return nil, err } - res := fmt.Sprintf(`OK`, account) + res := fmt.Sprintf(`OK`, EscapeXML(account)) lastDeviceID := "" for _, entry := range entries { @@ -234,13 +279,27 @@ func AccountFullToXML(ds *datastore.DataStore, account string) ([]byte, error) { continue } - res += fmt.Sprintf(``, deviceID) - res += fmt.Sprintf(`%s%s`, - info.ProductCode, info.ProductCode, info.ProductSerialNumber) + res += fmt.Sprintf(``, EscapeXML(deviceID)) + + res += fmt.Sprintf(``, EscapeXML(info.ProductCode)) + if len(info.Components) > 0 { + res += `` + for _, comp := range info.Components { + res += fmt.Sprintf(`%s%s%s`, + EscapeXML(comp.Category), EscapeXML(comp.Category), EscapeXML(comp.SoftwareVersion), EscapeXML(comp.SerialNumber)) + } + + res += `` + } else { + res += `` + } + + res += fmt.Sprintf(`%s%s`, + EscapeXML(info.ProductCode), EscapeXML(info.ProductSerialNumber)) res += fmt.Sprintf(`%s`, DateStr) - res += fmt.Sprintf(`%s`, info.FirmwareVersion) - res += fmt.Sprintf(`%s`, info.IPAddress) - res += fmt.Sprintf(`%s`, info.Name) + res += fmt.Sprintf(`%s`, EscapeXML(info.FirmwareVersion)) + res += fmt.Sprintf(`%s`, EscapeXML(info.IPAddress)) + res += fmt.Sprintf(`%s`, EscapeXML(info.Name)) presets, _ := PresetsToXML(ds, account, deviceID) if len(presets) > len(xml.Header) { @@ -338,12 +397,12 @@ func UpdatePreset(ds *datastore.DataStore, account, device string, presetNumber } // Return XML for the single preset - res := fmt.Sprintf(``, presetObj.ID) - res += fmt.Sprintf(`%s`, presetObj.ContainerArt) - res += fmt.Sprintf(`%s`, presetObj.Type) + res := fmt.Sprintf(``, EscapeXML(presetObj.ID)) + res += fmt.Sprintf(`%s`, EscapeXML(presetObj.ContainerArt)) + res += fmt.Sprintf(`%s`, EscapeXML(presetObj.Type)) res += fmt.Sprintf(`%s`, DateStr) - res += fmt.Sprintf(`%s`, presetObj.Location) - res += fmt.Sprintf(`%s`, presetObj.Name) + res += fmt.Sprintf(`%s`, EscapeXML(presetObj.Location)) + res += fmt.Sprintf(`%s`, EscapeXML(presetObj.Name)) res += GetConfiguredSourceXML(*matchingSrc) res += fmt.Sprintf(`%s`, DateStr) res += `` @@ -354,12 +413,12 @@ func UpdatePreset(ds *datastore.DataStore, account, device string, presetNumber // AddRecent adds or updates a recent item for the specified account and device. func AddRecent(ds *datastore.DataStore, account, device string, sourceXML []byte) ([]byte, error) { sources, err := ds.GetConfiguredSources(account, device) - if err != nil { + if err != nil && !os.IsNotExist(err) { return nil, err } recents, err := ds.GetRecents(account, device) - if err != nil { + if err != nil && !os.IsNotExist(err) { return nil, err } @@ -376,7 +435,25 @@ func AddRecent(ds *datastore.DataStore, account, device string, sourceXML []byte matchingSrc := findMatchingSource(sources, newRecentElem.SourceID) if matchingSrc == nil { - return nil, fmt.Errorf("invalid account/source") + // If we don't have a matching source, try to guess or create a virtual one. + // For Spotify, the location usually starts with /playback/container/c3... + // which is a base64 encoded spotify: URI. + if strings.Contains(newRecentElem.Location, "spotify") || newRecentElem.SourceID == "SPOTIFY" { + matchingSrc = &models.ConfiguredSource{ + ID: newRecentElem.SourceID, + DisplayName: "Spotify", + } + matchingSrc.SourceKey.Type = "SPOTIFY" + matchingSrc.SourceKeyType = "SPOTIFY" + } else { + // fallback to a generic source if we can't guess + matchingSrc = &models.ConfiguredSource{ + ID: newRecentElem.SourceID, + DisplayName: "Other", + } + matchingSrc.SourceKey.Type = "INVALID" + matchingSrc.SourceKeyType = "INVALID" + } } utcTime := parseLastPlayedAt(newRecentElem.LastPlayedAt) @@ -464,13 +541,14 @@ func createNewRecent(recents []models.ServiceRecent, name string, matchingSrc *m func formatRecentResponse(recentObj *models.ServiceRecent, matchingSrc *models.ConfiguredSource, createdOn string, utcTime int64) []byte { lastPlayed := time.Unix(utcTime, 0).Format(time.RFC3339) - res := fmt.Sprintf(``, recentObj.ID) - res += fmt.Sprintf(`%s`, recentObj.Type) - res += fmt.Sprintf(`%s`, createdOn) - res += fmt.Sprintf(`%s`, lastPlayed) - res += fmt.Sprintf(`%s`, recentObj.Location) - res += fmt.Sprintf(`%s`, recentObj.Name) + res := fmt.Sprintf(``, EscapeXML(recentObj.ID)) + res += fmt.Sprintf(`%s`, EscapeXML(recentObj.Type)) + res += fmt.Sprintf(`%s`, EscapeXML(createdOn)) + res += fmt.Sprintf(`%s`, EscapeXML(lastPlayed)) + res += fmt.Sprintf(`%s`, EscapeXML(recentObj.Location)) + res += fmt.Sprintf(`%s`, EscapeXML(recentObj.Name)) res += GetConfiguredSourceXML(*matchingSrc) + res += fmt.Sprintf(`%s`, EscapeXML(matchingSrc.ID)) res += fmt.Sprintf(`%s`, DateStr) res += `` @@ -498,11 +576,11 @@ func AddDeviceToAccount(ds *datastore.DataStore, account string, sourceXML []byt } createdOn := time.Now().Format(time.RFC3339) - res := fmt.Sprintf(``, newDeviceElem.DeviceID) - res += fmt.Sprintf(`%s`, createdOn) + res := fmt.Sprintf(``, EscapeXML(newDeviceElem.DeviceID)) + res += fmt.Sprintf(`%s`, EscapeXML(createdOn)) res += `` - res += fmt.Sprintf(`%s`, newDeviceElem.Name) - res += fmt.Sprintf(`%s`, createdOn) + res += fmt.Sprintf(`%s`, EscapeXML(newDeviceElem.Name)) + res += fmt.Sprintf(`%s`, EscapeXML(createdOn)) res += `` return append([]byte(xml.Header), []byte(res)...), nil diff --git a/pkg/service/marge/marge_test.go b/pkg/service/marge/marge_test.go index 7b90c61..0fb5808 100644 --- a/pkg/service/marge/marge_test.go +++ b/pkg/service/marge/marge_test.go @@ -2,6 +2,8 @@ package marge import ( "os" + "path/filepath" + "strconv" "strings" "testing" "time" @@ -64,6 +66,104 @@ func TestMargeXML(t *testing.T) { } } +func TestEscapeXML(t *testing.T) { + input := "Antenne Chillout & Other" + expected := "Antenne Chillout & Other" + actual := EscapeXML(input) + if actual != expected { + t.Errorf("Expected %s, got %s", expected, actual) + } + + inputWithAll := "< > & ' \"" + expectedWithAll := "< > & ' "" + actualWithAll := EscapeXML(inputWithAll) + if actualWithAll != expectedWithAll { + t.Errorf("Expected %s, got %s", expectedWithAll, actualWithAll) + } +} + +func TestRecentsXML_EmptyIDFix(t *testing.T) { + tempDir, err := os.MkdirTemp("", "marge-test-*") + if err != nil { + t.Fatalf("Failed to create temp dir: %v", err) + } + defer func() { _ = os.RemoveAll(tempDir) }() + + ds := datastore.NewDataStore(tempDir) + account := "test-acc" + device := "test-dev" + + deviceDir := ds.AccountDeviceDir(account, device) + _ = os.MkdirAll(deviceDir, 0755) + + // Create a Recents.xml with empty ID + recentsXML := []byte(` + + + + Test Item + + +`) + _ = os.WriteFile(filepath.Join(deviceDir, "Recents.xml"), recentsXML, 0644) + _ = os.WriteFile(filepath.Join(deviceDir, "Sources.xml"), []byte(""), 0644) + + // Fetching should fix the empty ID + recents, err := ds.GetRecents(account, device) + if err != nil { + t.Fatalf("Failed to get recents: %v", err) + } + + if len(recents) != 1 { + t.Fatalf("Expected 1 recent, got %d", len(recents)) + } + + if recents[0].ID == "" { + t.Errorf("Expected non-empty ID for recent") + } + + if _, err := strconv.Atoi(recents[0].ID); err != nil { + t.Errorf("Expected numeric ID, got %s", recents[0].ID) + } + + // Verify the XML output also has the non-empty ID + xmlData, err := RecentsToXML(ds, account, device) + if err != nil { + t.Fatalf("RecentsToXML failed: %v", err) + } + + if strings.Contains(string(xmlData), `recent id=""`) { + t.Errorf("XML should not contain empty recent ID: %s", string(xmlData)) + } + + if !strings.Contains(string(xmlData), `recent id="1"`) { + t.Errorf("XML should contain fixed numeric ID: %s", string(xmlData)) + } +} + +func TestGetConfiguredSourceXML_Escaping(t *testing.T) { + src := models.ConfiguredSource{ + ID: "101&202", + DisplayName: "Test & Source", + Secret: "key&value", + } + src.SourceKeyAccount = "user&name" + + xml := GetConfiguredSourceXML(src) + if !strings.Contains(xml, "id=\"101&202\"") { + t.Errorf("ID not escaped in attribute: %s", xml) + } + if strings.Contains(xml, "101&202") { + t.Errorf("ID should not be escaped in sourceid tag inside source tag anymore: %s", xml) + } + if !strings.Contains(xml, "Test & Source") { + t.Errorf("DisplayName not escaped: %s", xml) + } + if !strings.Contains(xml, ">key&value") { + t.Errorf("Secret not escaped: %s", xml) + } +} + func TestAddRecent_TimestampPreservation(t *testing.T) { tempDir, err := os.MkdirTemp("", "marge-test-*") if err != nil { @@ -112,9 +212,6 @@ func TestAddRecent_TimestampPreservation(t *testing.T) { t.Fatalf("Expected 1 recent, got %d", len(recents)) } - originalCreatedOn := recents[0].UtcTime // It's stored in UtcTime field (unix string) in models.ServiceRecent but the AddRecent return XML uses tag which is DateStr or Now depending on logic. - // Actually let's check what AddRecent returns. - // 3. Add the same recent again (it should move to front and preserve createdOn) // We'll wait a second to ensure time.Now() would be different if it were used for createdOn time.Sleep(1 * time.Second) @@ -134,9 +231,11 @@ func TestAddRecent_TimestampPreservation(t *testing.T) { t.Errorf("Expected still 1 recent, got %d", len(recents)) } - // Check that UtcTime was updated (it should be, for lastplayedat) - if recents[0].UtcTime == originalCreatedOn { - // Wait, if they are the same it might be because we didn't specify LastPlayedAt in input XML so it used Now. - // Since we slept, it should be different. + // Verify that sourceid is present in recent response and is a sibling to source tag + if !strings.Contains(string(respXML), "101") { + t.Errorf("Expected sourceid in recent response: %s", string(respXML)) + } + if strings.Contains(string(respXML), "2012-09-19T12:43:00.000+00:00key&valuetest-user101") { + t.Errorf("sourceid should not be inside source tag: %s", string(respXML)) } } diff --git a/pkg/service/proxy/recorder.go b/pkg/service/proxy/recorder.go index 090334e..4e40edc 100644 --- a/pkg/service/proxy/recorder.go +++ b/pkg/service/proxy/recorder.go @@ -117,9 +117,18 @@ func (r *Recorder) Record(category string, req *http.Request, res *http.Response // Clone request clonedReq = req.Clone(req.Context()) if req.Body != nil { - bodyBytes, _ := io.ReadAll(req.Body) - req.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) - clonedReq.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + bodyBytes, err := io.ReadAll(req.Body) + if err != nil { + log.Printf("failed to read request body for async recording: %v", err) + + clonedReq.Body = http.NoBody + } else { + // Reset original body for subsequent consumers (though Record is usually called at the end) + req.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + // Set body for async task + clonedReq.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + clonedReq.ContentLength = int64(len(bodyBytes)) + } } // Clone response if present @@ -130,9 +139,17 @@ func (r *Recorder) Record(category string, req *http.Request, res *http.Response Request: clonedReq, } if res.Body != nil { - bodyBytes, _ := io.ReadAll(res.Body) - res.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) - clonedRes.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + bodyBytes, err := io.ReadAll(res.Body) + if err != nil { + log.Printf("failed to read response body for async recording: %v", err) + + res.Body = http.NoBody + clonedRes.Body = http.NoBody + } else { + res.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + clonedRes.Body = io.NopCloser(bytes.NewBuffer(bodyBytes)) + clonedRes.ContentLength = int64(len(bodyBytes)) + } } } } else {