@@ -123,6 +123,9 @@ func (crt cacheRoundTripper) RoundTrip(req *http.Request) (*http.Response, error
123123// Allow an individual request to override cache options.
124124func requestCacheOptions (req * http.Request ) (string , time.Duration ) {
125125 var dur time.Duration
126+ // Added alongside the TTL header in https://github.com/cli/go-gh/pull/49.
127+ // No production consumer of the directory override is known; retain it for
128+ // compatibility.
126129 dir := req .Header .Get ("X-GH-CACHE-DIR" )
127130 ttl := req .Header .Get ("X-GH-CACHE-TTL" )
128131 if ttl != "" {
@@ -170,62 +173,91 @@ func (fs *fileStorage) read(key string) (*http.Response, error) {
170173 return res , err
171174}
172175
173- func (fs * fileStorage ) store (key string , res * http.Response ) (storeErr error ) {
174- cacheFile := fs .filePath (key )
175-
176+ func (fs * fileStorage ) store (key string , res * http.Response ) error {
176177 fs .mu .Lock ()
177178 defer fs .mu .Unlock ()
178179
179- dir := filepath .Dir (cacheFile )
180- if storeErr = os .MkdirAll (dir , 0755 ); storeErr != nil {
181- return
180+ cacheFilePath := fs .filePath (key )
181+ dir := filepath .Dir (cacheFilePath )
182+ if err := os .MkdirAll (dir , 0755 ); err != nil {
183+ return err
182184 }
183185
184- // Write to a temporary file in the same directory, then rename into place
185- // so concurrent processes and partial writes cannot corrupt the cache entry .
186- var f * os.File
187- if f , storeErr = os . CreateTemp ( dir , ".gh-cache-*" ); storeErr != nil {
188- return
186+ // Finish writing before publishing the entry. Same- directory rename gives
187+ // atomic replacement on Unix; it is best-effort on other platforms .
188+ tmpCacheFile , err := os .CreateTemp ( dir , ".gh-cache-*" )
189+ if err != nil {
190+ return err
189191 }
190- tmpName := f .Name ()
192+ tmpCacheFileName := tmpCacheFile .Name ()
193+ // Clean up on errors and panics too. After a successful rename, the
194+ // temporary path no longer exists, so removing it is harmless.
191195 defer func () {
192- if storeErr != nil {
193- _ = os .Remove (tmpName )
194- }
196+ _ = tmpCacheFile .Close ()
197+ _ = os .Remove (tmpCacheFileName )
195198 }()
196199
197- var origBody io.ReadCloser
198- if res .Body != nil {
199- origBody , res .Body = copyStream (res .Body )
200- defer res .Body .Close ()
200+ if err := writeCacheResponse (tmpCacheFile , res ); err != nil {
201+ return err
201202 }
202-
203- storeErr = res .Write (f )
204- if origBody != nil {
205- res .Body = origBody
203+ if err := tmpCacheFile .Close (); err != nil {
204+ return err
206205 }
207206
208- if cerr := f .Close (); storeErr == nil && cerr != nil {
209- storeErr = cerr
210- }
211- if storeErr != nil {
212- return
207+ return renameCacheFile (tmpCacheFileName , cacheFilePath )
208+ }
209+
210+ func writeCacheResponse (w io.Writer , res * http.Response ) error {
211+ if res .Body == nil {
212+ // Serialize the HTTP response headers only, since there is no body.
213+ return res .Write (w )
213214 }
214215
215- if err := os .Chmod (tmpName , 0600 ); err != nil {
216- storeErr = err
217- return
216+ // Buffer the bytes consumed during serialization so the caller can replay
217+ // them. Restore the replay reader even if writing fails or panics.
218+ buffer := & bytes.Buffer {}
219+ recorder := & errorRecordingReader {Reader : io .TeeReader (res .Body , buffer )}
220+ source := & readCloser {Reader : recorder , Closer : res .Body }
221+ res .Body = source
222+ defer source .Close ()
223+ defer func () {
224+ res .Body = io .NopCloser (& errorReplayingReader {Reader : buffer , err : recorder .err })
225+ }()
226+
227+ return res .Write (w )
228+ }
229+
230+ type errorRecordingReader struct {
231+ io.Reader
232+ err error
233+ }
234+
235+ func (r * errorRecordingReader ) Read (p []byte ) (int , error ) {
236+ n , err := r .Reader .Read (p )
237+ if err != nil && err != io .EOF {
238+ r .err = err
218239 }
240+ return n , err
241+ }
219242
220- storeErr = os .Rename (tmpName , cacheFile )
221- return
243+ type errorReplayingReader struct {
244+ io.Reader
245+ err error
246+ }
247+
248+ func (r * errorReplayingReader ) Read (p []byte ) (int , error ) {
249+ n , err := r .Reader .Read (p )
250+ if err == io .EOF && r .err != nil {
251+ err = r .err
252+ r .err = nil
253+ }
254+ return n , err
222255}
223256
224- func copyStream (r io.ReadCloser ) (io.ReadCloser , io.ReadCloser ) {
225- b := & bytes.Buffer {}
226- nr := io .TeeReader (r , b )
227- return io .NopCloser (b ), & readCloser {
228- Reader : nr ,
229- Closer : r ,
257+ func copyStream (body io.ReadCloser ) (replay , source io.ReadCloser ) {
258+ buffer := & bytes.Buffer {}
259+ return io .NopCloser (buffer ), & readCloser {
260+ Reader : io .TeeReader (body , buffer ),
261+ Closer : body ,
230262 }
231263}
0 commit comments