Add checks for deployment model name
This commit is contained in:
parent
577796ff13
commit
4101111df9
2 changed files with 34 additions and 3 deletions
|
|
@ -6,6 +6,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
|
"net/url"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -87,9 +88,13 @@ func (p *Provider) Chat(
|
||||||
// model is the deployment name for Azure OpenAI
|
// model is the deployment name for Azure OpenAI
|
||||||
deployment := model
|
deployment := model
|
||||||
|
|
||||||
// Build Azure-specific URL: {base}/openai/deployments/{deployment}/chat/completions?api-version=...
|
// Build Azure-specific URL safely using url.JoinPath and query encoding
|
||||||
requestURL := fmt.Sprintf("%s/openai/deployments/%s/chat/completions?api-version=%s",
|
// to prevent path traversal or query injection via deployment names.
|
||||||
p.apiBase, deployment, azureAPIVersion)
|
base, err := url.JoinPath(p.apiBase, "openai/deployments", deployment, "chat/completions")
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to build Azure request URL: %w", err)
|
||||||
|
}
|
||||||
|
requestURL := base + "?api-version=" + azureAPIVersion
|
||||||
|
|
||||||
// Build request body — no "model" field (Azure infers from deployment URL)
|
// Build request body — no "model" field (Azure infers from deployment URL)
|
||||||
requestBody := map[string]any{
|
requestBody := map[string]any{
|
||||||
|
|
|
||||||
|
|
@ -204,3 +204,29 @@ func TestProvider_AzureNewProviderWithTimeout(t *testing.T) {
|
||||||
t.Errorf("timeout = %v, want %v", p.httpClient.Timeout, 180*time.Second)
|
t.Errorf("timeout = %v, want %v", p.httpClient.Timeout, 180*time.Second)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProviderChat_AzureDeploymentNameEscaped(t *testing.T) {
|
||||||
|
var capturedPath string
|
||||||
|
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
capturedPath = r.URL.RawPath // use RawPath to see percent-encoding
|
||||||
|
if capturedPath == "" {
|
||||||
|
capturedPath = r.URL.Path
|
||||||
|
}
|
||||||
|
writeValidResponse(w)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
p := NewProvider("test-key", server.URL, "")
|
||||||
|
|
||||||
|
// Deployment name with characters that could cause path injection
|
||||||
|
_, err := p.Chat(t.Context(), []Message{{Role: "user", Content: "hi"}}, nil, "my deploy/../../admin", nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Chat() error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// The slash and special chars in the deployment name must be escaped, not treated as path separators
|
||||||
|
if capturedPath == "/openai/deployments/my deploy/../../admin/chat/completions" {
|
||||||
|
t.Fatal("deployment name was interpolated without escaping — path injection possible")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue