328 lines
9.9 KiB
Go
328 lines
9.9 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"`
|
|
}
|
|
|
|
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"`
|
|
}
|
|
|
|
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)
|
|
if request.Code == "" || len(request.HostFingerprint) < 16 ||
|
|
request.Hostname == "" || len(request.Hostname) > 255 {
|
|
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, "", sourceIP, "")
|
|
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
|
|
}
|
|
|
|
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 len(response.Packages) == 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), "")
|
|
}
|