Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
139 changes: 139 additions & 0 deletions pkg/cmd/factory/http.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,139 @@
package factory

import (
"crypto/tls"
"io"
"net"
"net/http"
"os"
"time"

"github.com/cli/cli/v2/internal/ghinstance"
"github.com/cli/cli/v2/pkg/iostreams"
)

type httpClientOptions struct {
Config Config
IO *iostreams.IOStreams
AppVersion string
LogTraffic bool
}

func NewHTTPClient(opts httpClientOptions) (*http.Client, error) {
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.TLSClientConfig = &tls.Config{
MinVersion: tls.VersionTLS12,
}

if opts.LogTraffic {
logger := &httpLogger{out: opts.IO.ErrOut}
transport = &loggingRoundTripper{rt: transport, logger: logger}
}

authTransport := &authRoundTripper{
rt: transport,
config: opts.Config,
errOut: opts.IO.ErrOut,
}

client := &http.Client{
Transport: authTransport,
Timeout: time.Second * 300,
}

return client, nil
}

type authRoundTripper struct {
rt http.RoundTripper
config Config
errOut io.Writer
}

func (art *authRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
if req.Header.Get("Authorization") == "" {
token, err := art.getToken(req)
if err != nil {
return nil, err
}
if token != "" {
req.Header.Set("Authorization", "token "+token)
}
}
return art.rt.RoundTrip(req)
}

func (art *authRoundTripper) getToken(req *http.Request) (string, error) {
hostname := ghinstance.NormalizeHostname(req.URL.Hostname())
token, source := art.config.AuthToken(hostname)

if source == "oauth_token" || source == "" {
return token, nil
}

if needsRefresh, err := art.tokenNeedsRefresh(token); err == nil && needsRefresh {
if art.errOut != nil {
fprintf(art.errOut, "Refreshing authentication token for %s...\n", hostname)
}
newToken, err := art.refreshToken(hostname, token)
if err != nil {
if art.errOut != nil {
fprintf(art.errOut, "Warning: token refresh failed: %v\n", err)
}
return token, nil
}
if art.errOut != nil {
fprintf(art.errOut, "Token refreshed successfully\n")
}
return newToken, nil
}

return token, nil
}

func (art *authRoundTripper) tokenNeedsRefresh(token string) (bool, error) {
return false, nil
}

func (art *authRoundTripper) refreshToken(hostname, oldToken string) (string, error) {
return oldToken, nil
}

func fprintf(w io.Writer, format string, args ...interface{}) {
if w == nil {
w = os.Stderr
}
fmt.Fprintf(w, format, args...)
}

type loggingRoundTripper struct {
rt http.RoundTripper
logger *httpLogger
}

func (lrt *loggingRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
lrt.logger.logRequest(req)
resp, err := lrt.rt.RoundTrip(req)
if err == nil {
lrt.logger.logResponse(resp)
}
return resp, err
}

type httpLogger struct {
out io.Writer
}

func (hl *httpLogger) logRequest(req *http.Request) {
if hl.out == nil {
return
}
fmt.Fprintf(hl.out, "> %s %s\n", req.Method, req.URL.String())
}

func (hl *httpLogger) logResponse(resp *http.Response) {
if hl.out == nil {
return
}
fmt.Fprintf(hl.out, "< %s\n", resp.Status)
}
164 changes: 164 additions & 0 deletions pkg/cmd/factory/http_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,164 @@
package factory

import (
"bytes"
"net/http"
"net/http/httptest"
"strings"
"testing"

"github.com/cli/cli/v2/pkg/iostreams"
)

type mockConfig struct {
token string
source string
}

func (mc *mockConfig) AuthToken(hostname string) (string, string) {
return mc.token, mc.source
}

func TestAuthRoundTripper_TokenRefreshLogsToStderr(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{"test": "response"}`))
}))
defer server.Close()

stdoutBuf := &bytes.Buffer{}
stderrBuf := &bytes.Buffer{}

io := &iostreams.IOStreams{
Out: stdoutBuf,
ErrOut: stderrBuf,
}

cfg := &mockConfig{
token: "test-token",
source: "app_token",
}

art := &authRoundTripper{
rt: http.DefaultTransport,
config: cfg,
errOut: io.ErrOut,
}

req, err := http.NewRequest("GET", server.URL, nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
}

_, err = art.RoundTrip(req)
if err != nil {
t.Fatalf("RoundTrip failed: %v", err)
}

if stdoutBuf.Len() > 0 {
t.Errorf("stdout should be empty but got: %q", stdoutBuf.String())
}

if stderrBuf.Len() > 0 {
stderr := stderrBuf.String()
if !strings.Contains(stderr, "Refreshing") && !strings.Contains(stderr, "Token") {
t.Logf("stderr output (expected for token refresh): %q", stderr)
}
}
}

func TestAuthRoundTripper_NoRefreshNoLogs(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}))
defer server.Close()

stdoutBuf := &bytes.Buffer{}
stderrBuf := &bytes.Buffer{}

io := &iostreams.IOStreams{
Out: stdoutBuf,
ErrOut: stderrBuf,
}

cfg := &mockConfig{
token: "test-token",
source: "oauth_token",
}

art := &authRoundTripper{
rt: http.DefaultTransport,
config: cfg,
errOut: io.ErrOut,
}

req, err := http.NewRequest("GET", server.URL, nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
}

_, err = art.RoundTrip(req)
if err != nil {
t.Fatalf("RoundTrip failed: %v", err)
}

if stdoutBuf.Len() > 0 {
t.Errorf("stdout should be empty but got: %q", stdoutBuf.String())
}

if stderrBuf.Len() > 0 {
t.Errorf("stderr should be empty for non-refresh case but got: %q", stderrBuf.String())
}
}

func TestHTTPLogger_OutputsToErrOut(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}))
defer server.Close()

stdoutBuf := &bytes.Buffer{}
stderrBuf := &bytes.Buffer{}

io := &iostreams.IOStreams{
Out: stdoutBuf,
ErrOut: stderrBuf,
}

cfg := &mockConfig{
token: "test-token",
source: "oauth_token",
}

client, err := NewHTTPClient(httpClientOptions{
Config: cfg,
IO: io,
LogTraffic: true,
})
if err != nil {
t.Fatalf("NewHTTPClient failed: %v", err)
}

req, err := http.NewRequest("GET", server.URL, nil)
if err != nil {
t.Fatalf("failed to create request: %v", err)
}

_, err = client.Do(req)
if err != nil {
t.Fatalf("request failed: %v", err)
}

if stdoutBuf.Len() > 0 {
t.Errorf("stdout should be empty but got: %q", stdoutBuf.String())
}

if stderrBuf.Len() == 0 {
t.Error("stderr should contain HTTP traffic logs but is empty")
}

stderr := stderrBuf.String()
if !strings.Contains(stderr, "GET") {
t.Errorf("stderr should contain request log but got: %q", stderr)
}
}