package proxy import ( "bytes" "compress/gzip" "crypto/sha256" "encoding/hex" "io" "net/http" "net/http/httptest" "os" "path/filepath" "strings" ) const cacheDir = "/tmp/ghpcache" func (h *ProxyHandler) serveWithCache(w http.ResponseWriter, r *http.Request, upstream http.Handler) { if !shouldConsiderForCache(r) { upstream.ServeHTTP(w, r) return } if body, ok := loadFromCache(r); ok { w.Header().Set("Content-Type", "text/html; charset=utf-8") w.WriteHeader(http.StatusOK) _, _ = w.Write(body) return } rec := httptest.NewRecorder() upstream.ServeHTTP(rec, r) resp := rec.Result() defer resp.Body.Close() rawBody, err := io.ReadAll(resp.Body) if err != nil { //NOTE(kroot): if we couldn't read the body, just proxy what we have in recorder for k, vals := range resp.Header { for _, v := range vals { w.Header().Add(k, v) } } w.WriteHeader(resp.StatusCode) _, _ = w.Write(rawBody) return } body := rawBody contentType := resp.Header.Get("Content-Type") encoding := strings.ToLower(resp.Header.Get("Content-Encoding")) //NOTE(kroot): if it's HTML and it's gzip, unpack before caching and serving if resp.StatusCode == http.StatusOK && isHTML(contentType) { if encoding == "gzip" { gzReader, gzErr := gzip.NewReader(bytes.NewReader(rawBody)) if gzErr == nil { defer gzReader.Close() if plainBody, readErr := io.ReadAll(gzReader); readErr == nil { body = plainBody // remove gzip encoding header resp.Header.Del("Content-Encoding") } } } _ = saveToCache(r, body) } for k, vals := range resp.Header { for _, v := range vals { w.Header().Add(k, v) } } //NOTE(kroot): if we changed the body (e.g. unpacked gzip), Content-Length may be incorrect w.Header().Del("Content-Length") w.WriteHeader(resp.StatusCode) _, _ = w.Write(body) } func shouldConsiderForCache(r *http.Request) bool { if r.Method != http.MethodGet { return false } path := r.URL.Path if strings.HasPrefix(path, "/api/") || strings.Contains(path, "/assets/") { return false } parts := strings.Split(strings.Trim(path, "/"), "/") if len(parts) == 2 { //NOTE(kroot): /owner/repo return true } return false } func cacheKey(r *http.Request) string { base := r.Host + "|" + r.URL.Path + "?" + r.URL.RawQuery sum := sha256.Sum256([]byte(base)) return hex.EncodeToString(sum[:]) } func cacheFilePath(r *http.Request) string { return filepath.Join(cacheDir, cacheKey(r)+".html") } func loadFromCache(r *http.Request) ([]byte, bool) { path := cacheFilePath(r) data, err := os.ReadFile(path) if err != nil { return nil, false } return data, true } func saveToCache(r *http.Request, body []byte) error { if err := os.MkdirAll(cacheDir, 0o755); err != nil { return err } path := cacheFilePath(r) return os.WriteFile(path, body, 0o644) } func isHTML(contentType string) bool { contentType = strings.ToLower(contentType) return strings.HasPrefix(contentType, "text/html") }