Files
goget/internal/hsts/hsts.go
T

227 lines
5.5 KiB
Go

//go:build linux || freebsd
// +build linux freebsd
// Package hsts implements HTTP Strict Transport Security (RFC 6797) caching.
package hsts
import (
"encoding/json"
"fmt"
"net/http"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"time"
"golang.org/x/net/publicsuffix"
)
// DefaultCachePath returns the default HSTS cache file path.
func DefaultCachePath() string {
if path := os.Getenv("GOGET_HSTS"); path != "" {
return path
}
if xdg := os.Getenv("XDG_CONFIG_HOME"); xdg != "" {
return filepath.Join(xdg, "goget", "hsts")
}
if home, err := os.UserHomeDir(); err == nil {
return filepath.Join(home, ".config", "goget", "hsts")
}
return ""
}
// Entry represents a single HSTS rule for a host.
type Entry struct {
Host string `json:"host"`
IncludeSubdomains bool `json:"include_subdomains"`
Expires time.Time `json:"expires"`
}
// Cache stores HSTS rules persistently.
type Cache struct {
entries map[string]*Entry
mu sync.RWMutex
path string
}
// NewCache creates an empty HSTS cache.
func NewCache() *Cache {
return &Cache{
entries: make(map[string]*Entry),
}
}
// Load reads HSTS entries from disk.
func (c *Cache) Load(path string) error {
c.mu.Lock()
defer c.mu.Unlock()
c.path = path
data, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return nil
}
return fmt.Errorf("failed to read hsts cache: %w", err)
}
var entries []*Entry
if err := json.Unmarshal(data, &entries); err != nil {
return fmt.Errorf("failed to parse hsts cache: %w", err)
}
now := time.Now()
for _, e := range entries {
if e.Expires.After(now) {
c.entries[e.Host] = e
}
}
return nil
}
// Save writes HSTS entries to disk.
func (c *Cache) Save() error {
c.mu.RLock()
defer c.mu.RUnlock()
if c.path == "" {
return nil
}
now := time.Now()
var active []*Entry
for _, e := range c.entries {
if e.Expires.After(now) {
active = append(active, e)
}
}
if len(active) == 0 {
_ = os.Remove(c.path)
return nil
}
data, err := json.MarshalIndent(active, "", " ")
if err != nil {
return fmt.Errorf("failed to marshal hsts cache: %w", err)
}
if err := os.MkdirAll(filepath.Dir(c.path), 0755); err != nil {
return fmt.Errorf("failed to create hsts directory: %w", err)
}
if err := os.WriteFile(c.path, data, 0644); err != nil {
return fmt.Errorf("failed to write hsts cache: %w", err)
}
return nil
}
// Apply upgrades http:// to https:// if the host is in the HSTS cache.
// Returns true if the URL was modified.
func (c *Cache) Apply(u *url.URL) bool {
if u.Scheme != "http" {
return false
}
c.mu.RLock()
defer c.mu.RUnlock()
host := strings.ToLower(u.Hostname())
if c.matchLocked(host) != nil {
u.Scheme = "https"
return true
}
return false
}
// Update parses Strict-Transport-Security from an HTTPS response and stores it.
// Per RFC 6797, HSTS headers on non-HTTPS responses are ignored.
func (c *Cache) Update(resp *http.Response) {
if resp == nil || resp.Request == nil || resp.Request.URL.Scheme != "https" {
return
}
header := resp.Header.Get("Strict-Transport-Security")
if header == "" {
return
}
maxAge, includeSubdomains := parseHeader(header)
if maxAge <= 0 {
return
}
host := strings.ToLower(resp.Request.URL.Hostname())
entry := &Entry{
Host: host,
IncludeSubdomains: includeSubdomains,
Expires: time.Now().Add(time.Duration(maxAge) * time.Second),
}
c.mu.Lock()
c.entries[host] = entry
c.mu.Unlock()
}
// matchLocked finds an entry for the given host (must hold read lock).
// When includeSubdomains is set on an entry, its rule applies to all
// subdomains, but only down to the registrable domain boundary
// (publicsuffix.EffectiveTLDPlusOne). Walking past the registrable
// domain would compare against a public suffix (e.g. "co.uk"), which
// RFC 6797 does not authorise and would let a HSTS entry mistakenly
// stored at the suffix level cover unrelated registrable domains.
func (c *Cache) matchLocked(host string) *Entry {
if e, ok := c.entries[host]; ok {
if e.Expires.After(time.Now()) {
return e
}
}
// Walk up parent domains, stopping once we step outside the
// registrable domain (eTLD+1). We check the public-suffix
// boundary *before* entries: a HSTS entry stored at a public
// suffix (e.g. "co.uk") must not cover unrelated registrable
// domains like "example.co.uk".
for {
dot := strings.Index(host, ".")
if dot < 0 {
break
}
host = host[dot+1:]
// If the parent is a public suffix or a single label
// (EffectiveTLDPlusOne returns an error), there is no further
// in-scope parent to consider.
if _, err := publicsuffix.EffectiveTLDPlusOne(host); err != nil {
break
}
if e, ok := c.entries[host]; ok && e.IncludeSubdomains && e.Expires.After(time.Now()) {
return e
}
}
return nil
}
// parseHeader parses the Strict-Transport-Security header value.
// Returns max-age in seconds and whether includeSubDomains is set.
func parseHeader(value string) (int64, bool) {
var maxAge int64 = -1
var includeSubdomains bool
for _, part := range strings.Split(value, ";") {
part = strings.TrimSpace(part)
if strings.HasPrefix(strings.ToLower(part), "max-age=") {
v := strings.TrimPrefix(part, "max-age=")
v = strings.TrimPrefix(v, "max-Age=")
v = strings.TrimPrefix(v, "MAX-AGE=")
if n, err := strconv.ParseInt(v, 10, 64); err == nil {
maxAge = n
}
} else if strings.EqualFold(part, "includesubdomains") {
includeSubdomains = true
}
}
return maxAge, includeSubdomains
}