301 lines
8.5 KiB
Go
301 lines
8.5 KiB
Go
package providerapi
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"net/url"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"proxy-pool/internal/config"
|
|
controllerProvider "proxy-pool/internal/controller/provider"
|
|
)
|
|
|
|
const defaultRetryAfterMax = 30 * time.Second
|
|
|
|
type HTTPDoer interface {
|
|
Do(*http.Request) (*http.Response, error)
|
|
}
|
|
|
|
type HTTPAdapter struct {
|
|
doer HTTPDoer
|
|
method string
|
|
url string
|
|
headers http.Header
|
|
body []byte
|
|
basicUser string
|
|
basicPass string
|
|
bearerToken string
|
|
maxBytes int64
|
|
retryAfterMax time.Duration
|
|
now func() time.Time
|
|
}
|
|
|
|
type HTTPStatusError struct {
|
|
StatusCode int
|
|
retryable bool
|
|
}
|
|
|
|
func (a *HTTPAdapter) String() string {
|
|
if a == nil {
|
|
return "providerapi.HTTPAdapter<nil>"
|
|
}
|
|
return fmt.Sprintf("providerapi.HTTPAdapter{method:%s,maxBytes:%d}", a.method, a.maxBytes)
|
|
}
|
|
|
|
func (a *HTTPAdapter) GoString() string { return a.String() }
|
|
|
|
func (e *HTTPStatusError) Error() string {
|
|
return fmt.Sprintf("provider API returned HTTP status %d", e.StatusCode)
|
|
}
|
|
|
|
func (e *HTTPStatusError) Retryable() bool { return e.retryable }
|
|
|
|
type operationError struct {
|
|
operation string
|
|
cause error
|
|
}
|
|
|
|
func (e *operationError) Error() string { return e.operation + " failed" }
|
|
func (e *operationError) Unwrap() error { return e.cause }
|
|
|
|
func NewHTTPAdapter(api config.ProviderAPI, fetch config.Fetch, doer HTTPDoer) (*HTTPAdapter, error) {
|
|
parsedURL, err := url.Parse(api.URL)
|
|
if err != nil || parsedURL.Host == "" || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: invalid API URL")
|
|
}
|
|
method := strings.ToUpper(strings.TrimSpace(api.Method))
|
|
if method == "" {
|
|
method = http.MethodGet
|
|
}
|
|
if !validHTTPToken(method) {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: invalid method %q", api.Method)
|
|
}
|
|
|
|
query := parsedURL.Query()
|
|
for key, value := range api.Query {
|
|
query.Set(key, value)
|
|
}
|
|
headers := make(http.Header, len(api.Headers)+2)
|
|
for key, value := range api.Headers {
|
|
if !validHTTPToken(key) || !validHTTPHeaderValue(value) {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: invalid header")
|
|
}
|
|
headers.Set(key, value)
|
|
}
|
|
|
|
adapter := &HTTPAdapter{
|
|
doer: doer,
|
|
method: method,
|
|
headers: headers,
|
|
maxBytes: fetch.MaxResponseBytes,
|
|
retryAfterMax: time.Duration(fetch.Retry.Max),
|
|
now: time.Now,
|
|
}
|
|
if adapter.doer == nil {
|
|
adapter.doer = defaultHTTPClient()
|
|
}
|
|
if adapter.maxBytes <= 0 {
|
|
adapter.maxBytes = defaultTemplateMaxBytes
|
|
}
|
|
if adapter.retryAfterMax <= 0 {
|
|
adapter.retryAfterMax = defaultRetryAfterMax
|
|
}
|
|
|
|
switch api.Auth.Type {
|
|
case "", "none":
|
|
case "basic":
|
|
if api.Auth.Username == "" || api.Auth.Password == "" {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: resolved basic credentials are required")
|
|
}
|
|
if headerHas(headers, "Authorization") {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: conflicting authorization header")
|
|
}
|
|
adapter.basicUser, adapter.basicPass = api.Auth.Username, api.Auth.Password
|
|
case "bearer":
|
|
if api.Auth.Token == "" {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: resolved bearer token is required")
|
|
}
|
|
if headerHas(headers, "Authorization") {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: conflicting authorization header")
|
|
}
|
|
adapter.bearerToken = api.Auth.Token
|
|
case "apiKey":
|
|
if api.Auth.Name == "" || api.Auth.Value == "" {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: resolved API key is required")
|
|
}
|
|
switch api.Auth.Location {
|
|
case "header":
|
|
if !validHTTPToken(api.Auth.Name) || !validHTTPHeaderValue(api.Auth.Value) {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: invalid API key header")
|
|
}
|
|
if headerHas(headers, api.Auth.Name) {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: conflicting API key header")
|
|
}
|
|
adapter.headers.Set(api.Auth.Name, api.Auth.Value)
|
|
case "query":
|
|
if query.Has(api.Auth.Name) {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: conflicting API key query parameter")
|
|
}
|
|
query.Set(api.Auth.Name, api.Auth.Value)
|
|
default:
|
|
return nil, fmt.Errorf("new provider HTTP adapter: invalid API key location %q", api.Auth.Location)
|
|
}
|
|
default:
|
|
return nil, fmt.Errorf("new provider HTTP adapter: unsupported auth type %q", api.Auth.Type)
|
|
}
|
|
parsedURL.RawQuery = query.Encode()
|
|
adapter.url = parsedURL.String()
|
|
|
|
switch api.Body.Type {
|
|
case "":
|
|
case "json":
|
|
adapter.body, err = json.Marshal(api.Body.Value)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("new provider HTTP adapter: encode JSON body: %w", err)
|
|
}
|
|
if adapter.headers.Get("Content-Type") == "" {
|
|
adapter.headers.Set("Content-Type", "application/json")
|
|
}
|
|
case "form":
|
|
values := make(url.Values, len(api.Body.Value))
|
|
for key, value := range api.Body.Value {
|
|
values.Set(key, value)
|
|
}
|
|
adapter.body = []byte(values.Encode())
|
|
if adapter.headers.Get("Content-Type") == "" {
|
|
adapter.headers.Set("Content-Type", "application/x-www-form-urlencoded")
|
|
}
|
|
default:
|
|
return nil, fmt.Errorf("new provider HTTP adapter: unsupported body type %q", api.Body.Type)
|
|
}
|
|
return adapter, nil
|
|
}
|
|
|
|
func (a *HTTPAdapter) Fetch(ctx context.Context) (controllerProvider.FetchResponse, error) {
|
|
request, err := http.NewRequestWithContext(ctx, a.method, a.url, bytes.NewReader(a.body))
|
|
if err != nil {
|
|
return controllerProvider.FetchResponse{}, &operationError{operation: "build provider request", cause: err}
|
|
}
|
|
request.Header = a.headers.Clone()
|
|
if a.basicUser != "" {
|
|
request.SetBasicAuth(a.basicUser, a.basicPass)
|
|
}
|
|
if a.bearerToken != "" {
|
|
request.Header.Set("Authorization", "Bearer "+a.bearerToken)
|
|
}
|
|
|
|
response, err := a.doer.Do(request)
|
|
if err != nil {
|
|
if response != nil && response.Body != nil {
|
|
_ = response.Body.Close()
|
|
}
|
|
return controllerProvider.FetchResponse{}, &operationError{operation: "call provider API", cause: err}
|
|
}
|
|
if response == nil || response.Body == nil {
|
|
return controllerProvider.FetchResponse{}, ErrInvalidHTTPResponse
|
|
}
|
|
defer response.Body.Close()
|
|
body, readErr := readAllLimited(response.Body, a.maxBytes, ErrResponseTooLarge)
|
|
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
|
|
statusErr := &HTTPStatusError{
|
|
StatusCode: response.StatusCode,
|
|
retryable: retryableHTTPStatus(response.StatusCode),
|
|
}
|
|
result := controllerProvider.FetchResponse{}
|
|
if statusErr.retryable {
|
|
result.RetryAfter = parseRetryAfter(response.Header.Get("Retry-After"), a.now(), a.retryAfterMax)
|
|
}
|
|
if readErr != nil {
|
|
if !errors.Is(readErr, ErrResponseTooLarge) {
|
|
readErr = &operationError{operation: "read provider response", cause: readErr}
|
|
}
|
|
return result, errors.Join(statusErr, readErr)
|
|
}
|
|
return result, statusErr
|
|
}
|
|
if readErr != nil {
|
|
if errors.Is(readErr, ErrResponseTooLarge) {
|
|
return controllerProvider.FetchResponse{}, readErr
|
|
}
|
|
return controllerProvider.FetchResponse{}, &operationError{operation: "read provider response", cause: readErr}
|
|
}
|
|
return controllerProvider.FetchResponse{Body: body}, nil
|
|
}
|
|
|
|
func retryableHTTPStatus(status int) bool {
|
|
return status == http.StatusRequestTimeout || status == http.StatusTooEarly ||
|
|
status == http.StatusTooManyRequests ||
|
|
(status >= http.StatusInternalServerError && status <= 599)
|
|
}
|
|
|
|
func parseRetryAfter(value string, now time.Time, maximum time.Duration) time.Duration {
|
|
value = strings.TrimSpace(value)
|
|
if value == "" {
|
|
return 0
|
|
}
|
|
if seconds, err := strconv.ParseInt(value, 10, 64); err == nil {
|
|
if seconds <= 0 {
|
|
return 0
|
|
}
|
|
if seconds > int64(maximum/time.Second) {
|
|
return maximum
|
|
}
|
|
return time.Duration(seconds) * time.Second
|
|
}
|
|
when, err := http.ParseTime(value)
|
|
if err != nil {
|
|
return 0
|
|
}
|
|
delay := when.Sub(now)
|
|
if delay <= 0 {
|
|
return 0
|
|
}
|
|
if delay > maximum {
|
|
return maximum
|
|
}
|
|
return delay
|
|
}
|
|
|
|
func defaultHTTPClient() *http.Client {
|
|
return &http.Client{
|
|
CheckRedirect: func(*http.Request, []*http.Request) error {
|
|
return http.ErrUseLastResponse
|
|
},
|
|
}
|
|
}
|
|
|
|
func validHTTPToken(value string) bool {
|
|
if value == "" {
|
|
return false
|
|
}
|
|
for _, character := range value {
|
|
if character <= 0x20 || character >= 0x7f || strings.ContainsRune("()<>@,;:\\\"/[]?={}", character) {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func validHTTPHeaderValue(value string) bool {
|
|
for _, character := range value {
|
|
if character == '\t' {
|
|
continue
|
|
}
|
|
if character < 0x20 || character == 0x7f {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func headerHas(headers http.Header, key string) bool {
|
|
_, exists := headers[http.CanonicalHeaderKey(key)]
|
|
return exists
|
|
}
|