110 lines
3.4 KiB
Go
110 lines
3.4 KiB
Go
package processing
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"net/http"
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
"time"
|
||
|
|
)
|
||
|
|
|
||
|
|
func TestNewOpenAIClient_capsRetries(t *testing.T) {
|
||
|
|
c := NewOpenAIClient("k", "https://api.openai.com/v1", "m", 0, 99)
|
||
|
|
if c.MaxRetries != maxOpenAIRetries {
|
||
|
|
t.Fatalf("MaxRetries=%d want %d", c.MaxRetries, maxOpenAIRetries)
|
||
|
|
}
|
||
|
|
c2 := NewOpenAIClient("k", "", "m", 0, 0)
|
||
|
|
if c2.MaxRetries != 3 {
|
||
|
|
t.Fatalf("default MaxRetries=%d want 3", c2.MaxRetries)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestNewOpenAIClient_blocksPrivateDial(t *testing.T) {
|
||
|
|
c := NewOpenAIClient("k", "https://api.openai.com/v1", "m", 0, 1)
|
||
|
|
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://127.0.0.1:9/", nil)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||
|
|
defer cancel()
|
||
|
|
req = req.WithContext(ctx)
|
||
|
|
_, err = c.HTTPClient.Do(req)
|
||
|
|
if err == nil {
|
||
|
|
t.Fatal("expected dial to private/loopback blocked")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestOpenAIBaseAllowsLoopback(t *testing.T) {
|
||
|
|
if !openAIBaseAllowsLoopback("http://localhost:11434/v1") {
|
||
|
|
t.Fatal("expected localhost allowed")
|
||
|
|
}
|
||
|
|
if !openAIBaseAllowsLoopback("http://127.0.0.1:11434/v1") {
|
||
|
|
t.Fatal("expected 127.0.0.1 allowed")
|
||
|
|
}
|
||
|
|
if openAIBaseAllowsLoopback("https://api.openai.com/v1") {
|
||
|
|
t.Fatal("expected public host denied for loopback flag")
|
||
|
|
}
|
||
|
|
if openAIBaseAllowsLoopback("https://192.168.1.1/v1") {
|
||
|
|
t.Fatal("expected private IP denied")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestOpenAIBaseAllowsPrivateNonProd(t *testing.T) {
|
||
|
|
t.Setenv("APP_ENV", "local")
|
||
|
|
if !openAIBaseAllowsPrivate("http://192.168.50.181:8767/v1") {
|
||
|
|
t.Fatal("expected LAN proxy allowed in local")
|
||
|
|
}
|
||
|
|
if !openAIDialPolicy("http://192.168.50.181:8767/v1").AllowPrivate {
|
||
|
|
t.Fatal("expected dial policy AllowPrivate")
|
||
|
|
}
|
||
|
|
t.Setenv("APP_ENV", "production")
|
||
|
|
if openAIBaseAllowsPrivate("http://192.168.50.181:8767/v1") {
|
||
|
|
t.Fatal("expected LAN proxy blocked in production")
|
||
|
|
}
|
||
|
|
t.Setenv("APP_ENV", "local")
|
||
|
|
if openAIBaseAllowsPrivate("http://169.254.169.254/v1") {
|
||
|
|
t.Fatal("expected link-local metadata blocked")
|
||
|
|
}
|
||
|
|
if openAIBaseAllowsPrivate("https://api.openai.com/v1") {
|
||
|
|
t.Fatal("expected public host not private-allowed")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestNewOpenAIClient_allowsPrivateDialNonProd(t *testing.T) {
|
||
|
|
t.Setenv("APP_ENV", "development")
|
||
|
|
c := NewOpenAIClient("k", "http://192.168.50.181:8767/v1", "m", 0, 1)
|
||
|
|
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://192.168.50.181:9/", nil)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||
|
|
defer cancel()
|
||
|
|
req = req.WithContext(ctx)
|
||
|
|
_, err = c.HTTPClient.Do(req)
|
||
|
|
if err == nil {
|
||
|
|
t.Fatal("expected connection error, not success")
|
||
|
|
}
|
||
|
|
if strings.Contains(err.Error(), "host is not allowed") {
|
||
|
|
t.Fatalf("SSRF blocked LAN OpenAI base unexpectedly: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestNewOpenAIClient_blocksPrivateDialInProduction(t *testing.T) {
|
||
|
|
t.Setenv("APP_ENV", "production")
|
||
|
|
c := NewOpenAIClient("k", "http://192.168.50.181:8767/v1", "m", 0, 1)
|
||
|
|
req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "http://192.168.50.181:9/", nil)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||
|
|
defer cancel()
|
||
|
|
req = req.WithContext(ctx)
|
||
|
|
_, err = c.HTTPClient.Do(req)
|
||
|
|
if err == nil {
|
||
|
|
t.Fatal("expected dial blocked in production")
|
||
|
|
}
|
||
|
|
if !strings.Contains(err.Error(), "host is not allowed") {
|
||
|
|
t.Fatalf("expected host is not allowed, got: %v", err)
|
||
|
|
}
|
||
|
|
}
|