diff options
Diffstat (limited to 'web/server.go')
| -rw-r--r-- | web/server.go | 108 |
1 files changed, 89 insertions, 19 deletions
diff --git a/web/server.go b/web/server.go index cc3ad2e..53374ce 100644 --- a/web/server.go +++ b/web/server.go @@ -4,6 +4,7 @@ import ( "crypto/rand" "database/sql" "encoding/hex" + "encoding/json" "errors" "fmt" "log/slog" @@ -32,6 +33,8 @@ type Server struct { oauthMu sync.Mutex oauthState string returnToURL string + syncMu sync.Mutex + syncActive bool } func NewServer(db *sql.DB, configPath string) *Server { @@ -139,7 +142,7 @@ func (s *Server) handleSyncForm(w http.ResponseWriter, r *http.Request) { } func (s *Server) handleSync(w http.ResponseWriter, r *http.Request) { - if err := r.ParseForm(); err != nil { + if err := r.ParseMultipartForm(10 << 20); err != nil && !errors.Is(err, http.ErrNotMultipart) { http.Error(w, "invalid form", http.StatusBadRequest) return } @@ -164,6 +167,11 @@ func (s *Server) handleSync(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "/strava/login?"+query.Encode(), http.StatusFound) return } + if !s.beginSync() { + http.Error(w, "a Strava sync is already in progress", http.StatusConflict) + return + } + defer s.endSync() if refreshed, err := client.RefreshIfNeeded(); err != nil { s.renderSyncError(w, r, r.FormValue("from"), r.FormValue("to"), err) @@ -174,17 +182,22 @@ func (s *Server) handleSync(w http.ResponseWriter, r *http.Request) { return } } - imported, skipped, err := s.syncRides(client, appConfig.Storage.GPXDir, from, to) + flusher, ok := w.(http.Flusher) + if !ok { + http.Error(w, "streaming responses are not supported", http.StatusInternalServerError) + return + } + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Content-Type", "text/event-stream") + progress, err := s.syncRides(client, appConfig.Storage.GPXDir, from, to, func(progress SyncProgress) error { + return writeSyncEvent(w, flusher, "progress", progress) + }) if err != nil { - s.renderSyncError(w, r, r.FormValue("from"), r.FormValue("to"), err) + _ = writeSyncEvent(w, flusher, "error", syncErrorEvent{Message: err.Error(), Progress: progress}) return } - query := url.Values{} - query.Set("from", r.FormValue("from")) - query.Set("to", r.FormValue("to")) - query.Set("imported", strconv.Itoa(imported)) - query.Set("skipped", strconv.Itoa(skipped)) - http.Redirect(w, r, "/sync?"+query.Encode(), http.StatusSeeOther) + slog.Info("Completed Strava sync", "from", from.Format(dateFormat), "to", to.Add(-24*time.Hour).Format(dateFormat), "total", progress.Total, "imported", progress.Imported, "skipped", progress.Skipped) + _ = writeSyncEvent(w, flusher, "complete", progress) } func (s *Server) handleStravaLogin(w http.ResponseWriter, r *http.Request) { @@ -250,31 +263,40 @@ func (s *Server) handleStravaCallback(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, returnTo, http.StatusSeeOther) } -func (s *Server) syncRides(client *strava.Client, gpxDir string, from, to time.Time) (imported, skipped int, err error) { +func (s *Server) syncRides(client *strava.Client, gpxDir string, from, to time.Time, report func(SyncProgress) error) (SyncProgress, error) { activities, err := client.List(from, to) if err != nil { - return 0, 0, err + return SyncProgress{}, err + } + progress := SyncProgress{Total: len(activities)} + if err := reportSyncProgress(report, progress); err != nil { + return progress, err } + slog.Info("Started Strava sync", "from", from.Format(dateFormat), "to", to.Add(-24*time.Hour).Format(dateFormat), "total", progress.Total) if err := os.MkdirAll(gpxDir, 0o755); err != nil { - return 0, 0, fmt.Errorf("create GPX directory: %w", err) + return progress, fmt.Errorf("create GPX directory: %w", err) } for _, summary := range activities { externalID := fmt.Sprintf("strava:%d", summary.ID) _, exists, err := rides.GetByExternalID(s.db, externalID) if err != nil { - return imported, skipped, err + return progress, err } if exists { - skipped++ + progress.Skipped++ + progress.Completed++ + if err := reportSyncProgress(report, progress); err != nil { + return progress, err + } continue } activity, gpxData, err := client.Get(summary.ID) if err != nil { - return imported, skipped, err + return progress, err } gpxPath := filepath.Join(gpxDir, fmt.Sprintf("activity_%d.gpx", activity.ID)) if err := os.WriteFile(gpxPath, gpxData, 0o644); err != nil { - return imported, skipped, fmt.Errorf("write GPX for activity %d: %w", activity.ID, err) + return progress, fmt.Errorf("write GPX for activity %d: %w", activity.ID, err) } activityType := activity.SportType if activityType == "" { @@ -292,12 +314,56 @@ func (s *Server) syncRides(client *strava.Client, gpxDir string, from, to time.T TotalElevationGainM: activity.TotalElevationGainM, AverageSpeedMps: activity.AverageSpeedMps, }); err != nil { - return imported, skipped, fmt.Errorf("save activity %d: %w", activity.ID, err) + return progress, fmt.Errorf("save activity %d: %w", activity.ID, err) + } + progress.Imported++ + progress.Completed++ + if err := reportSyncProgress(report, progress); err != nil { + return progress, err } - imported++ slog.Info("Imported Strava ride", "activity", activity.ID, "name", activity.Name) } - return imported, skipped, nil + return progress, nil +} + +func reportSyncProgress(report func(SyncProgress) error, progress SyncProgress) error { + if report == nil { + return nil + } + return report(progress) +} + +func (s *Server) beginSync() bool { + s.syncMu.Lock() + defer s.syncMu.Unlock() + if s.syncActive { + return false + } + s.syncActive = true + return true +} + +func (s *Server) endSync() { + s.syncMu.Lock() + s.syncActive = false + s.syncMu.Unlock() +} + +type syncErrorEvent struct { + Message string `json:"message"` + Progress SyncProgress `json:"progress"` +} + +func writeSyncEvent(w http.ResponseWriter, flusher http.Flusher, event string, data any) error { + payload, err := json.Marshal(data) + if err != nil { + return err + } + if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", event, payload); err != nil { + return err + } + flusher.Flush() + return nil } func (s *Server) stravaClient() (*strava.Client, bool, error) { @@ -360,6 +426,10 @@ func (s *Server) consumeOAuthState(state string) (string, bool) { } func (s *Server) renderSyncError(w http.ResponseWriter, r *http.Request, from, to string, err error) { + if strings.Contains(r.Header.Get("Accept"), "text/event-stream") { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } w.WriteHeader(http.StatusBadRequest) kcore.RenderPage(r.Context(), SyncPage(SyncPageData{From: from, To: to, Error: err.Error(), HasAuth: s.hasStravaToken()}), w) } |