Files
TA-Deployment-Broker/internal/app/download.go
T
David Schroeder 0ed0fb5817 update
2026-07-26 10:33:50 -05:00

416 lines
12 KiB
Go

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"`
}
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)
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,
)
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 > UTC_TIMESTAMP(6) - INTERVAL 10 MINUTE`,
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 > UTC_TIMESTAMP(6)
FOR UPDATE`,
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 = UTC_TIMESTAMP(6)
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 > UTC_TIMESTAMP(6)
AND a.revoked_at IS NULL
AND a.expires_at > UTC_TIMESTAMP(6)
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), "")
}