package app import ( "database/sql" "encoding/json" "errors" "fmt" "io" "net/http" "net/url" "strings" "time" ) type exchangeRequest struct { Code string `json:"code"` HostFingerprint string `json:"host_fingerprint"` Hostname string `json:"hostname"` RequestedAction string `json:"requested_action,omitempty"` RequestedPackage string `json:"requested_package,omitempty"` InstallationID string `json:"installation_id,omitempty"` } type exchangePackage struct { Slug string `json:"slug"` DisplayName string `json:"display_name"` Version string `json:"version"` SHA256 string `json:"sha256"` DownloadURL string `json:"download_url"` } type exchangeResponse struct { SessionToken string `json:"session_token"` ExpiresAt time.Time `json:"expires_at"` Packages []exchangePackage `json:"packages"` Actions []string `json:"actions"` } func (s *Server) handleExchange(w http.ResponseWriter, r *http.Request) { sourceIP := s.clientIP(r) limited, err := s.exchangeRateLimited(r, sourceIP) if err != nil { writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "unable to validate request"}) return } if limited { w.Header().Set("Retry-After", "600") writeJSON(w, http.StatusTooManyRequests, map[string]string{"error": "too many failed attempts; try again later"}) return } r.Body = http.MaxBytesReader(w, r.Body, 32<<10) var request exchangeRequest if err := json.NewDecoder(r.Body).Decode(&request); err != nil { writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid request"}) return } request.Code = normalizeCode(request.Code) request.HostFingerprint = strings.TrimSpace(request.HostFingerprint) request.Hostname = strings.TrimSpace(request.Hostname) request.RequestedAction = strings.TrimSpace(request.RequestedAction) request.RequestedPackage = strings.TrimSpace(request.RequestedPackage) request.InstallationID = strings.TrimSpace(request.InstallationID) if request.Code == "" || len(request.HostFingerprint) < 16 || request.Hostname == "" || len(request.Hostname) > 255 || (request.RequestedAction != "" && !validSlug(request.RequestedAction)) || (request.RequestedPackage != "" && !validSlug(request.RequestedPackage)) { writeJSON(w, http.StatusBadRequest, map[string]string{"error": "code and host identity are required"}) return } response, authorizationID, err := s.exchangeDeploymentCode(r, request) if err != nil { status := http.StatusInternalServerError message := "unable to create download session" if errors.Is(err, errAuthorizationDenied) { status = http.StatusForbidden message = "authorization is invalid, expired, revoked, or at its host limit" } _ = s.audit(r.Context(), "code_exchange_failed", "", nil, request.Hostname, "", sourceIP, err.Error()) writeJSON(w, status, map[string]string{"error": message}) return } _ = s.audit( r.Context(), "code_exchanged", "", &authorizationID, request.Hostname, request.RequestedPackage, sourceIP, "requested_action="+request.RequestedAction, ) s.verifyFleetInstallation(r.Context(), request.InstallationID) writeJSON(w, http.StatusOK, response) } func (s *Server) exchangeRateLimited(r *http.Request, sourceIP string) (bool, error) { var attempts int err := s.db.QueryRowContext( r.Context(), `SELECT COUNT(*) FROM audit_events WHERE event_type = 'code_exchange_failed' AND source_ip = ? AND created_at > datetime('now', '-10 minutes')`, sourceIP, ).Scan(&attempts) return attempts >= 10, err } func (s *Server) exchangeDeploymentCode( r *http.Request, request exchangeRequest, ) (exchangeResponse, uint64, error) { var response exchangeResponse codeHash := hashValue(request.Code) fingerprintHash := hashValue(request.HostFingerprint) tx, err := s.db.BeginTx(r.Context(), &sql.TxOptions{Isolation: sql.LevelSerializable}) if err != nil { return response, 0, err } defer tx.Rollback() var authorizationID uint64 var hostLimit int var expiresAt time.Time err = tx.QueryRowContext( r.Context(), `SELECT id, host_limit, expires_at FROM authorizations WHERE code_hash = ? AND revoked_at IS NULL AND expires_at > CURRENT_TIMESTAMP`, codeHash[:], ).Scan(&authorizationID, &hostLimit, &expiresAt) if errors.Is(err, sql.ErrNoRows) { return response, 0, errAuthorizationDenied } if err != nil { return response, 0, err } if request.RequestedAction != "" { var authorized int if err := tx.QueryRowContext( r.Context(), `SELECT COUNT(*) FROM authorization_actions aa JOIN installer_actions ia ON ia.slug = aa.action_slug WHERE aa.authorization_id = ? AND aa.action_slug = ? AND ia.enabled = TRUE`, authorizationID, request.RequestedAction, ).Scan(&authorized); err != nil { return response, 0, err } if authorized != 1 { return response, 0, errAuthorizationDenied } } if request.RequestedPackage != "" { var authorized int if err := tx.QueryRowContext( r.Context(), `SELECT COUNT(*) FROM authorization_packages ap JOIN packages p ON p.id = ap.package_id WHERE ap.authorization_id = ? AND p.slug = ? AND p.enabled = TRUE`, authorizationID, request.RequestedPackage, ).Scan(&authorized); err != nil { return response, 0, err } if authorized != 1 { return response, 0, errAuthorizationDenied } } var hostID uint64 err = tx.QueryRowContext( r.Context(), `SELECT id FROM authorization_hosts WHERE authorization_id = ? AND host_fingerprint = ?`, authorizationID, fingerprintHash[:], ).Scan(&hostID) if errors.Is(err, sql.ErrNoRows) { var hostCount int if err := tx.QueryRowContext( r.Context(), `SELECT COUNT(*) FROM authorization_hosts WHERE authorization_id = ?`, authorizationID, ).Scan(&hostCount); err != nil { return response, 0, err } if hostCount >= hostLimit { return response, 0, errAuthorizationDenied } result, err := tx.ExecContext( r.Context(), `INSERT INTO authorization_hosts (authorization_id, host_fingerprint, hostname) VALUES (?, ?, ?)`, authorizationID, fingerprintHash[:], request.Hostname, ) if err != nil { return response, 0, err } insertedID, _ := result.LastInsertId() hostID = uint64(insertedID) } else if err != nil { return response, 0, err } else { _, err = tx.ExecContext( r.Context(), `UPDATE authorization_hosts SET hostname = ?, last_seen_at = CURRENT_TIMESTAMP WHERE id = ?`, request.Hostname, hostID, ) if err != nil { return response, 0, err } } sessionToken, err := randomToken(32) if err != nil { return response, 0, err } sessionHash := hashValue(sessionToken) _, err = tx.ExecContext( r.Context(), `INSERT INTO download_sessions (token_hash, authorization_id, authorization_host_id, expires_at) VALUES (?, ?, ?, ?)`, sessionHash[:], authorizationID, hostID, expiresAt, ) if err != nil { return response, 0, err } rows, err := tx.QueryContext( r.Context(), `SELECT p.slug, p.display_name, p.package_version, p.sha256 FROM authorization_packages ap JOIN packages p ON p.id = ap.package_id WHERE ap.authorization_id = ? AND p.enabled = TRUE ORDER BY p.display_name`, authorizationID, ) if err != nil { return response, 0, err } defer rows.Close() for rows.Next() { var record exchangePackage if err := rows.Scan( &record.Slug, &record.DisplayName, &record.Version, &record.SHA256, ); err != nil { return response, 0, err } record.DownloadURL = strings.TrimRight(s.cfg.PublicURL.String(), "/") + "/api/v1/packages/" + url.PathEscape(record.Slug) response.Packages = append(response.Packages, record) } if err := rows.Err(); err != nil { return response, 0, err } if err := rows.Close(); err != nil { return response, 0, err } actionRows, err := tx.QueryContext( r.Context(), `SELECT ia.slug FROM authorization_actions aa JOIN installer_actions ia ON ia.slug = aa.action_slug WHERE aa.authorization_id = ? AND ia.enabled = TRUE ORDER BY ia.sort_order, ia.display_name`, authorizationID, ) if err != nil { return response, 0, err } defer actionRows.Close() for actionRows.Next() { var slug string if err := actionRows.Scan(&slug); err != nil { return response, 0, err } response.Actions = append(response.Actions, slug) } if err := actionRows.Err(); err != nil { return response, 0, err } if err := actionRows.Close(); err != nil { return response, 0, err } if len(response.Packages) == 0 && len(response.Actions) == 0 { return response, 0, errAuthorizationDenied } if err := tx.Commit(); err != nil { return response, 0, err } response.SessionToken = sessionToken response.ExpiresAt = expiresAt return response, authorizationID, nil } func (s *Server) handlePackageDownload(w http.ResponseWriter, r *http.Request) { const bearerPrefix = "Bearer " authorization := r.Header.Get("Authorization") if !strings.HasPrefix(authorization, bearerPrefix) { writeJSON(w, http.StatusUnauthorized, map[string]string{"error": "bearer token required"}) return } sessionToken := strings.TrimSpace(strings.TrimPrefix(authorization, bearerPrefix)) if sessionToken == "" { writeJSON(w, http.StatusUnauthorized, map[string]string{"error": "bearer token required"}) return } slug := r.PathValue("slug") sessionHash := hashValue(sessionToken) var packageInfo packageRecord var authorizationID uint64 var hostname string err := s.db.QueryRowContext( r.Context(), `SELECT p.id, p.slug, p.display_name, p.package_name, p.package_version, p.file_name, p.sha256, p.enabled, ds.authorization_id, ah.hostname FROM download_sessions ds JOIN authorizations a ON a.id = ds.authorization_id JOIN authorization_hosts ah ON ah.id = ds.authorization_host_id JOIN authorization_packages ap ON ap.authorization_id = a.id JOIN packages p ON p.id = ap.package_id WHERE ds.token_hash = ? AND ds.revoked_at IS NULL AND ds.expires_at > CURRENT_TIMESTAMP AND a.revoked_at IS NULL AND a.expires_at > CURRENT_TIMESTAMP AND p.slug = ? AND p.enabled = TRUE`, sessionHash[:], slug, ).Scan( &packageInfo.ID, &packageInfo.Slug, &packageInfo.DisplayName, &packageInfo.PackageName, &packageInfo.PackageVersion, &packageInfo.FileName, &packageInfo.SHA256, &packageInfo.Enabled, &authorizationID, &hostname, ) if errors.Is(err, sql.ErrNoRows) { writeJSON(w, http.StatusForbidden, map[string]string{"error": "package is not authorized"}) return } if err != nil { writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "unable to authorize package"}) return } registryURL := joinURL( s.cfg.GiteaURL, fmt.Sprintf( "/api/packages/%s/generic/%s/%s/%s", url.PathEscape(s.cfg.GiteaPackageOwner), url.PathEscape(packageInfo.PackageName), url.PathEscape(packageInfo.PackageVersion), url.PathEscape(packageInfo.FileName), ), ) upstreamRequest, err := http.NewRequestWithContext(r.Context(), http.MethodGet, registryURL, nil) if err != nil { writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "unable to request package"}) return } upstreamRequest.SetBasicAuth(s.cfg.GiteaPackageUser, s.cfg.GiteaPackageToken) upstreamResponse, err := s.packageClient.Do(upstreamRequest) if err != nil { _ = s.audit(r.Context(), "package_download_failed", "", &authorizationID, hostname, slug, s.clientIP(r), err.Error()) writeJSON(w, http.StatusBadGateway, map[string]string{"error": "package registry is unavailable"}) return } defer upstreamResponse.Body.Close() if upstreamResponse.StatusCode != http.StatusOK { _ = s.audit(r.Context(), "package_download_failed", "", &authorizationID, hostname, slug, s.clientIP(r), upstreamResponse.Status) writeJSON(w, http.StatusBadGateway, map[string]string{"error": "package registry rejected the request"}) return } w.Header().Set("Content-Type", "application/octet-stream") w.Header().Set("Content-Disposition", fmt.Sprintf("attachment; filename=%q", packageInfo.FileName)) w.Header().Set("X-TAPM-SHA256", packageInfo.SHA256) w.Header().Set("Cache-Control", "no-store") w.WriteHeader(http.StatusOK) if _, err := io.Copy(w, upstreamResponse.Body); err != nil { _ = s.audit(r.Context(), "package_download_failed", "", &authorizationID, hostname, slug, s.clientIP(r), err.Error()) return } _ = s.audit(r.Context(), "package_downloaded", "", &authorizationID, hostname, slug, s.clientIP(r), "") }