proxy-pool/internal/adapters/providerapi/http_adapter.go
youfak 4de3ffb85f
Some checks are pending
ci / test (ubuntu-latest) (push) Waiting to run
ci / test (windows-latest) (push) Waiting to run
ci / race (push) Waiting to run
feat: add ephemeral proxy activity pool
2026-07-29 12:51:18 +08:00

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
}