341 lines
10 KiB
Go
341 lines
10 KiB
Go
package providerapi
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"net/url"
|
|
"regexp"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"text/template"
|
|
"time"
|
|
|
|
"proxy-pool/internal/config"
|
|
proxyDomain "proxy-pool/internal/domain/proxy"
|
|
"proxy-pool/internal/platform/credentials"
|
|
)
|
|
|
|
const (
|
|
defaultTemplateTimeout = 100 * time.Millisecond
|
|
defaultTemplateMaxBytes = int64(1 << 20)
|
|
defaultMaxCandidates = 10_000
|
|
maxRegexPatterns = 64
|
|
maxRegexPatternBytes = 1024
|
|
)
|
|
|
|
type TemplateParser struct {
|
|
upstreamID string
|
|
template *template.Template
|
|
defaultScheme proxyDomain.Scheme
|
|
allowed map[proxyDomain.Scheme]struct{}
|
|
proxyAuthType string
|
|
username string
|
|
password string
|
|
maxConcurrency int64
|
|
timeout time.Duration
|
|
maxInputBytes int64
|
|
maxOutputBytes int64
|
|
maxCandidates int
|
|
credentialStore credentials.Store
|
|
executions *executionLimiter
|
|
|
|
regexMu sync.Mutex
|
|
regexes map[string]*regexp.Regexp
|
|
}
|
|
|
|
func (p *TemplateParser) String() string {
|
|
if p == nil {
|
|
return "providerapi.TemplateParser<nil>"
|
|
}
|
|
return fmt.Sprintf(
|
|
"providerapi.TemplateParser{upstream:%q,maxInputBytes:%d,maxOutputBytes:%d,maxCandidates:%d}",
|
|
p.upstreamID,
|
|
p.maxInputBytes,
|
|
p.maxOutputBytes,
|
|
p.maxCandidates,
|
|
)
|
|
}
|
|
|
|
func (p *TemplateParser) GoString() string { return p.String() }
|
|
|
|
func NewTemplateParser(
|
|
upstreamID string,
|
|
upstream config.Upstream,
|
|
credentialStore credentials.Store,
|
|
) (*TemplateParser, error) {
|
|
if strings.TrimSpace(upstreamID) == "" {
|
|
return nil, fmt.Errorf("new provider template parser: upstream ID is required")
|
|
}
|
|
parser := &TemplateParser{
|
|
upstreamID: upstreamID,
|
|
allowed: make(map[proxyDomain.Scheme]struct{}),
|
|
proxyAuthType: upstream.ProxyAuth.Type,
|
|
username: upstream.ProxyAuth.Username,
|
|
password: upstream.ProxyAuth.Password,
|
|
maxConcurrency: int64(upstream.Capacity.MaxConcurrencyPerProxy),
|
|
timeout: time.Duration(upstream.Fetch.TemplateTimeout),
|
|
maxInputBytes: upstream.Fetch.MaxResponseBytes,
|
|
maxOutputBytes: upstream.Fetch.MaxResponseBytes,
|
|
maxCandidates: upstream.Pool.MaxSize,
|
|
credentialStore: credentialStore,
|
|
executions: newExecutionLimiter(upstream.Fetch.MaxInFlight),
|
|
regexes: make(map[string]*regexp.Regexp),
|
|
}
|
|
if parser.timeout <= 0 {
|
|
parser.timeout = defaultTemplateTimeout
|
|
}
|
|
if parser.maxInputBytes <= 0 {
|
|
parser.maxInputBytes = defaultTemplateMaxBytes
|
|
}
|
|
if parser.maxOutputBytes <= 0 {
|
|
parser.maxOutputBytes = defaultTemplateMaxBytes
|
|
}
|
|
if parser.maxCandidates <= 0 {
|
|
parser.maxCandidates = defaultMaxCandidates
|
|
}
|
|
if parser.maxConcurrency <= 0 {
|
|
parser.maxConcurrency = 1
|
|
}
|
|
if parser.proxyAuthType == "" {
|
|
parser.proxyAuthType = "response"
|
|
}
|
|
switch parser.proxyAuthType {
|
|
case "response", "static", "ipWhitelist":
|
|
default:
|
|
return nil, fmt.Errorf("new provider template parser: unsupported proxy auth type %q", parser.proxyAuthType)
|
|
}
|
|
for _, protocol := range upstream.Provider.Protocols {
|
|
scheme := proxyDomain.Scheme(strings.ToLower(strings.TrimSpace(protocol)))
|
|
if !supportedScheme(scheme) {
|
|
return nil, fmt.Errorf("new provider template parser: unsupported protocol %q", protocol)
|
|
}
|
|
if parser.defaultScheme == "" {
|
|
parser.defaultScheme = scheme
|
|
}
|
|
parser.allowed[scheme] = struct{}{}
|
|
}
|
|
if parser.defaultScheme == "" {
|
|
parser.defaultScheme = proxyDomain.SchemeHTTP
|
|
parser.allowed[parser.defaultScheme] = struct{}{}
|
|
}
|
|
|
|
parsed, err := template.New("provider-response").Option("missingkey=error").Funcs(template.FuncMap{
|
|
"regexFind": parser.regexFind,
|
|
"regexFindAll": parser.regexFindAll,
|
|
}).Parse(upstream.API.Template)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("new provider template parser: parse template: %w", err)
|
|
}
|
|
if err := validateTemplateComplexity(parsed); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := rejectRecursiveTemplates(parsed); err != nil {
|
|
return nil, err
|
|
}
|
|
parser.template = parsed
|
|
return parser, nil
|
|
}
|
|
|
|
func (p *TemplateParser) Parse(ctx context.Context, body []byte) ([]proxyDomain.Proxy, error) {
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := enforceByteLimit(ErrTemplateInputTooLarge, int64(len(body)), p.maxInputBytes); err != nil {
|
|
return nil, err
|
|
}
|
|
execCtx, cancel := context.WithTimeout(ctx, p.timeout)
|
|
defer cancel()
|
|
output := limitedBuffer{ctx: execCtx, kind: ErrTemplateOutputTooLarge, limit: p.maxOutputBytes}
|
|
err := p.executions.Run(execCtx, func() error {
|
|
return p.template.Execute(&output, string(body))
|
|
})
|
|
if ctxErr := execCtx.Err(); ctxErr != nil {
|
|
return nil, ctxErr
|
|
}
|
|
if err != nil {
|
|
return nil, fmt.Errorf("execute provider template: %w", err)
|
|
}
|
|
|
|
tokens := strings.Fields(output.String())
|
|
if len(tokens) > p.maxCandidates {
|
|
return nil, &limitError{kind: ErrTooManyCandidates, size: int64(len(tokens)), limit: int64(p.maxCandidates)}
|
|
}
|
|
proxies := make([]proxyDomain.Proxy, 0, len(tokens))
|
|
credentialIndexes := make(map[string]int)
|
|
for _, token := range tokens {
|
|
candidate, credential, ok := p.parseCandidate(token)
|
|
if ok {
|
|
credentialKey := ""
|
|
if credential != nil {
|
|
if p.credentialStore == nil {
|
|
return nil, ErrCredentialStoreRequired
|
|
}
|
|
reference, err := p.credentialStore.Put(ctx, p.credentialScope(candidate), *credential)
|
|
if err != nil {
|
|
return nil, &operationError{operation: "store provider credentials", cause: err}
|
|
}
|
|
if reference.SecretRef == "" || reference.CredentialVersion == "" {
|
|
return nil, &operationError{
|
|
operation: "store provider credentials",
|
|
cause: credentials.ErrInvalidReference,
|
|
}
|
|
}
|
|
candidate.SecretRef = reference.SecretRef
|
|
candidate.CredentialVersion = reference.CredentialVersion
|
|
credentialKey = candidateCredentialKey(candidate)
|
|
if index, exists := credentialIndexes[credentialKey]; exists {
|
|
proxies[index] = candidate
|
|
continue
|
|
}
|
|
}
|
|
if len(proxies) >= p.maxCandidates {
|
|
return nil, &limitError{kind: ErrTooManyCandidates, size: int64(len(proxies) + 1), limit: int64(p.maxCandidates)}
|
|
}
|
|
proxies = append(proxies, candidate)
|
|
if credentialKey != "" {
|
|
credentialIndexes[credentialKey] = len(proxies) - 1
|
|
}
|
|
}
|
|
}
|
|
if len(tokens) > 0 && len(proxies) == 0 {
|
|
return nil, ErrInvalidProxyOutput
|
|
}
|
|
return proxies, nil
|
|
}
|
|
|
|
func (p *TemplateParser) regexFind(pattern, value string) (string, error) {
|
|
compiled, err := p.compileRegex(pattern)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
return compiled.FindString(value), nil
|
|
}
|
|
|
|
func (p *TemplateParser) regexFindAll(pattern, value string, count int) ([]string, error) {
|
|
compiled, err := p.compileRegex(pattern)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
matchLimit := p.maxCandidates
|
|
if matchLimit < int(^uint(0)>>1) {
|
|
matchLimit++
|
|
}
|
|
effectiveCount := count
|
|
if effectiveCount < 0 || effectiveCount > matchLimit {
|
|
effectiveCount = matchLimit
|
|
}
|
|
matches := compiled.FindAllString(value, effectiveCount)
|
|
if len(matches) > p.maxCandidates {
|
|
return nil, &limitError{kind: ErrTooManyCandidates, size: int64(len(matches)), limit: int64(p.maxCandidates)}
|
|
}
|
|
return matches, nil
|
|
}
|
|
|
|
func (p *TemplateParser) compileRegex(pattern string) (*regexp.Regexp, error) {
|
|
if len(pattern) > maxRegexPatternBytes {
|
|
return nil, fmt.Errorf("regex pattern exceeds %d bytes", maxRegexPatternBytes)
|
|
}
|
|
p.regexMu.Lock()
|
|
defer p.regexMu.Unlock()
|
|
if compiled := p.regexes[pattern]; compiled != nil {
|
|
return compiled, nil
|
|
}
|
|
if len(p.regexes) >= maxRegexPatterns {
|
|
return nil, fmt.Errorf("template exceeds %d distinct regex patterns", maxRegexPatterns)
|
|
}
|
|
compiled, err := regexp.Compile(pattern)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("compile template regex: invalid pattern")
|
|
}
|
|
p.regexes[pattern] = compiled
|
|
return compiled, nil
|
|
}
|
|
|
|
func (p *TemplateParser) parseCandidate(raw string) (proxyDomain.Proxy, *credentials.Value, bool) {
|
|
if !strings.Contains(raw, "://") {
|
|
raw = string(p.defaultScheme) + "://" + raw
|
|
}
|
|
parsed, err := url.Parse(raw)
|
|
if err != nil || parsed.Opaque != "" || parsed.Hostname() == "" || parsed.Port() == "" ||
|
|
parsed.Path != "" || parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" {
|
|
return proxyDomain.Proxy{}, nil, false
|
|
}
|
|
scheme := proxyDomain.Scheme(strings.ToLower(parsed.Scheme))
|
|
if _, ok := p.allowed[scheme]; !ok || !supportedScheme(scheme) {
|
|
return proxyDomain.Proxy{}, nil, false
|
|
}
|
|
port, err := strconv.ParseUint(parsed.Port(), 10, 16)
|
|
if err != nil || port == 0 {
|
|
return proxyDomain.Proxy{}, nil, false
|
|
}
|
|
username := ""
|
|
var credential *credentials.Value
|
|
switch p.proxyAuthType {
|
|
case "static":
|
|
username = p.username
|
|
if p.username != "" || p.password != "" {
|
|
credential = &credentials.Value{Username: p.username, Password: p.password}
|
|
}
|
|
case "response":
|
|
if parsed.User != nil {
|
|
username = parsed.User.Username()
|
|
password, hasPassword := parsed.User.Password()
|
|
if username != "" || hasPassword {
|
|
credential = &credentials.Value{Username: username, Password: password}
|
|
}
|
|
}
|
|
}
|
|
return proxyDomain.Proxy{
|
|
Scheme: scheme,
|
|
Host: parsed.Hostname(),
|
|
Port: uint16(port),
|
|
Username: username,
|
|
SourceUpstream: p.upstreamID,
|
|
MaxConcurrency: p.maxConcurrency,
|
|
State: proxyDomain.StateFetched,
|
|
}, credential, true
|
|
}
|
|
|
|
func (p *TemplateParser) credentialScope(candidate proxyDomain.Proxy) string {
|
|
if p.proxyAuthType == "static" {
|
|
return lengthPrefixedScope("provider", p.upstreamID, "static")
|
|
}
|
|
return lengthPrefixedScope(
|
|
"provider",
|
|
p.upstreamID,
|
|
string(candidate.Scheme),
|
|
strings.ToLower(candidate.Host),
|
|
strconv.FormatUint(uint64(candidate.Port), 10),
|
|
candidate.Username,
|
|
)
|
|
}
|
|
|
|
func lengthPrefixedScope(parts ...string) string {
|
|
var scope strings.Builder
|
|
for _, part := range parts {
|
|
scope.WriteString(strconv.Itoa(len(part)))
|
|
scope.WriteByte(':')
|
|
scope.WriteString(part)
|
|
}
|
|
return scope.String()
|
|
}
|
|
|
|
func candidateCredentialKey(candidate proxyDomain.Proxy) string {
|
|
return lengthPrefixedScope(
|
|
string(candidate.Scheme),
|
|
strings.ToLower(candidate.Host),
|
|
strconv.FormatUint(uint64(candidate.Port), 10),
|
|
candidate.Username,
|
|
)
|
|
}
|
|
|
|
func supportedScheme(scheme proxyDomain.Scheme) bool {
|
|
switch scheme {
|
|
case proxyDomain.SchemeHTTP, proxyDomain.SchemeHTTPS, proxyDomain.SchemeSOCKS5:
|
|
return true
|
|
default:
|
|
return false
|
|
}
|
|
}
|