added multi-key support for optional auth logic
This commit is contained in:
@@ -76,9 +76,10 @@ ollama-llama3,-,http://localhost:11434/v1,
|
||||
```
|
||||
|
||||
### 3. Setup gateway key (optional)
|
||||
To restrict access to your gateway, write a single plaintext key to a file (e.g. `gatekey.txt`):
|
||||
To restrict access to your gateway, write one or more plaintext keys (one per line) to a file (e.g. `gatekey.txt`):
|
||||
```text
|
||||
my-secure-gateway-token
|
||||
my-secure-gateway-token-1
|
||||
my-secure-gateway-token-2
|
||||
```
|
||||
|
||||
### 4. Run the server
|
||||
@@ -94,7 +95,7 @@ my-secure-gateway-token
|
||||
| ---- | ---- | ------- | ----------- |
|
||||
| `-port` | int | `8080` | Port the gateway listens on. |
|
||||
| `-csv` | string | `models.csv` | Path to the model configuration CSV. |
|
||||
| `-key` | string | `""` | Path to a file containing the gateway's access token. If blank, authentication is disabled. |
|
||||
| `-key` | string | `""` | Path to a file containing the gateway's access tokens (one per line). If blank, authentication is disabled. |
|
||||
| `-csv-updater` | string | `""` | Command/script to run periodically to update `models.csv`. |
|
||||
| `-csv-update-interval`| int | `10` | Frequency in minutes to invoke `-csv-updater`. |
|
||||
|
||||
@@ -152,7 +153,7 @@ curl -i -X POST http://localhost:8080/v1/chat/completions \
|
||||
|
||||
### 401 Unauthorized errors
|
||||
- **Symptom**: Requests return `{"error": {"code": "invalid_api_key", ...}}`.
|
||||
- **Solution**: Check that you are passing the correct Bearer token configured in the `-key` token file. Trim any whitespace, and verify it starts with `Bearer `.
|
||||
- **Solution**: Check that you are passing one of the correct Bearer tokens configured in the `-key` token file. Trim any whitespace, and verify it starts with `Bearer `.
|
||||
|
||||
## FAQ
|
||||
|
||||
|
||||
+108
-14
@@ -211,7 +211,7 @@ func TestGatewayCompletionsFailover(t *testing.T) {
|
||||
uniqueModels: []string{"model-1", "model-2"},
|
||||
}
|
||||
|
||||
handler := handleChatCompletions(cm, "")
|
||||
handler := handleChatCompletions(cm, nil)
|
||||
|
||||
// 1. Test non-streaming failover request starting with model-1
|
||||
reqBodyObj := map[string]any{
|
||||
@@ -284,7 +284,7 @@ func TestGatewayCompletionsStreamingFailover(t *testing.T) {
|
||||
uniqueModels: []string{"model-1", "model-2"},
|
||||
}
|
||||
|
||||
handler := handleChatCompletions(cm, "")
|
||||
handler := handleChatCompletions(cm, nil)
|
||||
|
||||
reqBodyObj := map[string]any{
|
||||
"model": "model-1",
|
||||
@@ -328,7 +328,7 @@ func TestModelsEndpoint(t *testing.T) {
|
||||
uniqueModels: []string{"model-1", "model-2"},
|
||||
}
|
||||
|
||||
handler := handleModels(cm, "")
|
||||
handler := handleModels(cm, nil)
|
||||
|
||||
req := httptest.NewRequest("GET", "/v1/models", nil)
|
||||
w := httptest.NewRecorder()
|
||||
@@ -367,9 +367,10 @@ func TestGatewayAuthentication(t *testing.T) {
|
||||
uniqueModels: []string{"model-1"},
|
||||
}
|
||||
|
||||
// Set up handlers expecting "secret-token"
|
||||
modelsHandler := handleModels(cm, "secret-token")
|
||||
completionsHandler := handleChatCompletions(cm, "secret-token")
|
||||
// Set up handlers expecting multiple secret tokens
|
||||
expectedTokens := []string{"secret-token-1", "secret-token-2"}
|
||||
modelsHandler := handleModels(cm, expectedTokens)
|
||||
completionsHandler := handleChatCompletions(cm, expectedTokens)
|
||||
|
||||
// Case 1: No Authorization header
|
||||
{
|
||||
@@ -392,14 +393,25 @@ func TestGatewayAuthentication(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// Case 3: Correct Authorization header
|
||||
// Case 3: Correct Authorization header (first token)
|
||||
{
|
||||
req := httptest.NewRequest("GET", "/v1/models", nil)
|
||||
req.Header.Set("Authorization", "Bearer secret-token")
|
||||
req.Header.Set("Authorization", "Bearer secret-token-1")
|
||||
w := httptest.NewRecorder()
|
||||
modelsHandler.ServeHTTP(w, req)
|
||||
if w.Result().StatusCode != http.StatusOK {
|
||||
t.Errorf("Expected 200 OK for correct token, got %d", w.Result().StatusCode)
|
||||
t.Errorf("Expected 200 OK for correct token 1, got %d", w.Result().StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
// Case 3b: Correct Authorization header (second token)
|
||||
{
|
||||
req := httptest.NewRequest("GET", "/v1/models", nil)
|
||||
req.Header.Set("Authorization", "Bearer secret-token-2")
|
||||
w := httptest.NewRecorder()
|
||||
modelsHandler.ServeHTTP(w, req)
|
||||
if w.Result().StatusCode != http.StatusOK {
|
||||
t.Errorf("Expected 200 OK for correct token 2, got %d", w.Result().StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -410,7 +422,7 @@ func TestGatewayAuthentication(t *testing.T) {
|
||||
}
|
||||
reqBytes, _ := json.Marshal(reqBodyObj)
|
||||
req := httptest.NewRequest("POST", "/v1/chat/completions", bytes.NewReader(reqBytes))
|
||||
req.Header.Set("Authorization", "Bearer secret-token")
|
||||
req.Header.Set("Authorization", "Bearer secret-token-2")
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
completionsHandler.ServeHTTP(w, req)
|
||||
@@ -492,7 +504,7 @@ func TestGatewayImageGenerations(t *testing.T) {
|
||||
uniqueModels: []string{"dall-e-3"},
|
||||
}
|
||||
|
||||
handler := handleImageGenerations(cm, "")
|
||||
handler := handleImageGenerations(cm, nil)
|
||||
|
||||
reqBodyObj := map[string]any{
|
||||
"prompt": "a beautiful kitten",
|
||||
@@ -551,7 +563,7 @@ func TestGatewayNoAuthHeaderOnEmptyKey(t *testing.T) {
|
||||
uniqueModels: []string{"model-no-key", "model-no-key-2"},
|
||||
}
|
||||
|
||||
handler := handleChatCompletions(cm, "")
|
||||
handler := handleChatCompletions(cm, nil)
|
||||
|
||||
// Case 1: Empty Key
|
||||
{
|
||||
@@ -615,7 +627,7 @@ func TestGatewayBlankKey(t *testing.T) {
|
||||
uniqueModels: []string{"model-blank-key"},
|
||||
}
|
||||
|
||||
handler := handleChatCompletions(cm, "")
|
||||
handler := handleChatCompletions(cm, nil)
|
||||
|
||||
reqBodyObj := map[string]any{
|
||||
"model": "model-blank-key",
|
||||
@@ -664,7 +676,7 @@ func TestGatewayExtraHeaders(t *testing.T) {
|
||||
uniqueModels: []string{"model-extra-headers"},
|
||||
}
|
||||
|
||||
handler := handleChatCompletions(cm, "")
|
||||
handler := handleChatCompletions(cm, nil)
|
||||
|
||||
reqBodyObj := map[string]any{
|
||||
"model": "model-extra-headers",
|
||||
@@ -685,5 +697,87 @@ func TestGatewayExtraHeaders(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestMultipleTokensKeyFile(t *testing.T) {
|
||||
// Create a temporary key file with multiple tokens, some empty lines, and comments/spaces
|
||||
tmpDir, err := os.MkdirTemp("", "dynagate-keyfile-test")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create temp dir: %v", err)
|
||||
}
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
keyFilePath := filepath.Join(tmpDir, "keys.txt")
|
||||
keyFileContent := "\n token-a \n\ntoken-b\n \ntoken-c\n"
|
||||
if err := os.WriteFile(keyFilePath, []byte(keyFileContent), 0644); err != nil {
|
||||
t.Fatalf("failed to write key file: %v", err)
|
||||
}
|
||||
|
||||
// Parse using the same logic as main.go
|
||||
content, err := os.ReadFile(keyFilePath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read key file: %v", err)
|
||||
}
|
||||
var authTokens []string
|
||||
for _, line := range strings.Split(string(content), "\n") {
|
||||
token := strings.TrimSpace(line)
|
||||
if token != "" {
|
||||
authTokens = append(authTokens, token)
|
||||
}
|
||||
}
|
||||
|
||||
expected := []string{"token-a", "token-b", "token-c"}
|
||||
if !reflect.DeepEqual(authTokens, expected) {
|
||||
t.Errorf("Expected parsed tokens to be %v, got %v", expected, authTokens)
|
||||
}
|
||||
|
||||
// Verify that the helper checkAuth works with these parsed tokens
|
||||
req := httptest.NewRequest("GET", "/v1/models", nil)
|
||||
req.Header.Set("Authorization", "Bearer token-b")
|
||||
if !checkAuth(authTokens, req) {
|
||||
t.Errorf("Expected token-b to authenticate successfully")
|
||||
}
|
||||
|
||||
reqWrong := httptest.NewRequest("GET", "/v1/models", nil)
|
||||
reqWrong.Header.Set("Authorization", "Bearer token-wrong")
|
||||
if checkAuth(authTokens, reqWrong) {
|
||||
t.Errorf("Expected token-wrong to fail authentication")
|
||||
}
|
||||
|
||||
// Test empty key file content behavior
|
||||
emptyKeyFilePath := filepath.Join(tmpDir, "empty_keys.txt")
|
||||
emptyKeyFileContent := "\n \n\n \n"
|
||||
if err := os.WriteFile(emptyKeyFilePath, []byte(emptyKeyFileContent), 0644); err != nil {
|
||||
t.Fatalf("failed to write empty key file: %v", err)
|
||||
}
|
||||
|
||||
// Parse using empty key file logic as in main.go
|
||||
emptyContent, err := os.ReadFile(emptyKeyFilePath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read empty key file: %v", err)
|
||||
}
|
||||
var emptyAuthTokens []string
|
||||
var tokens []string
|
||||
for _, line := range strings.Split(string(emptyContent), "\n") {
|
||||
token := strings.TrimSpace(line)
|
||||
if token != "" {
|
||||
tokens = append(tokens, token)
|
||||
}
|
||||
}
|
||||
if len(tokens) > 0 {
|
||||
emptyAuthTokens = tokens
|
||||
} else {
|
||||
emptyAuthTokens = nil
|
||||
}
|
||||
|
||||
if emptyAuthTokens != nil {
|
||||
t.Errorf("Expected emptyAuthTokens to be nil, got %v", emptyAuthTokens)
|
||||
}
|
||||
|
||||
// Verify that checkAuth returns true when expectedTokens is nil
|
||||
reqNoAuth := httptest.NewRequest("GET", "/v1/models", nil)
|
||||
if !checkAuth(emptyAuthTokens, reqNoAuth) {
|
||||
t.Errorf("Expected checkAuth to return true for nil expectedTokens")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
+14
-9
@@ -14,8 +14,8 @@ import (
|
||||
|
||||
var httpClient = &http.Client{}
|
||||
|
||||
func checkAuth(expectedToken string, r *http.Request) bool {
|
||||
if expectedToken == "" {
|
||||
func checkAuth(expectedTokens []string, r *http.Request) bool {
|
||||
if expectedTokens == nil {
|
||||
return true
|
||||
}
|
||||
authHeader := r.Header.Get("Authorization")
|
||||
@@ -23,7 +23,12 @@ func checkAuth(expectedToken string, r *http.Request) bool {
|
||||
return false
|
||||
}
|
||||
token := strings.TrimPrefix(authHeader, "Bearer ")
|
||||
return token == expectedToken
|
||||
for _, expected := range expectedTokens {
|
||||
if token == expected {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func sendUnauthorized(w http.ResponseWriter) {
|
||||
@@ -39,7 +44,7 @@ func sendUnauthorized(w http.ResponseWriter) {
|
||||
})
|
||||
}
|
||||
|
||||
func handleModels(cm *ConfigManager, expectedToken string) http.HandlerFunc {
|
||||
func handleModels(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
|
||||
type ModelData struct {
|
||||
ID string `json:"id"`
|
||||
Object string `json:"object"`
|
||||
@@ -53,7 +58,7 @@ func handleModels(cm *ConfigManager, expectedToken string) http.HandlerFunc {
|
||||
}
|
||||
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
if !checkAuth(expectedToken, r) {
|
||||
if !checkAuth(expectedTokens, r) {
|
||||
sendUnauthorized(w)
|
||||
return
|
||||
}
|
||||
@@ -91,9 +96,9 @@ func handleModels(cm *ConfigManager, expectedToken string) http.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
func handleChatCompletions(cm *ConfigManager, expectedToken string) http.HandlerFunc {
|
||||
func handleChatCompletions(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
if !checkAuth(expectedToken, r) {
|
||||
if !checkAuth(expectedTokens, r) {
|
||||
sendUnauthorized(w)
|
||||
return
|
||||
}
|
||||
@@ -324,9 +329,9 @@ func handleChatCompletions(cm *ConfigManager, expectedToken string) http.Handler
|
||||
}
|
||||
}
|
||||
|
||||
func handleImageGenerations(cm *ConfigManager, expectedToken string) http.HandlerFunc {
|
||||
func handleImageGenerations(cm *ConfigManager, expectedTokens []string) http.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request) {
|
||||
if !checkAuth(expectedToken, r) {
|
||||
if !checkAuth(expectedTokens, r) {
|
||||
sendUnauthorized(w)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -27,14 +27,25 @@ func main() {
|
||||
csvUpdateInterval := flag.Int("csv-update-interval", 10, "Interval in minutes with which to run the CSV updater command")
|
||||
flag.Parse()
|
||||
|
||||
var authToken string
|
||||
var authTokens []string
|
||||
if *keyPath != "" {
|
||||
content, err := os.ReadFile(*keyPath)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to read auth token file: %v", err)
|
||||
}
|
||||
authToken = strings.TrimSpace(string(content))
|
||||
log.Printf("Starting Dynagate LLM Gateway. Config: %s, Port: %d, Auth: Required", *csvPath, *port)
|
||||
var tokens []string
|
||||
for _, line := range strings.Split(string(content), "\n") {
|
||||
token := strings.TrimSpace(line)
|
||||
if token != "" {
|
||||
tokens = append(tokens, token)
|
||||
}
|
||||
}
|
||||
if len(tokens) > 0 {
|
||||
authTokens = tokens
|
||||
log.Printf("Starting Dynagate LLM Gateway. Config: %s, Port: %d, Auth: Required (%d tokens loaded)", *csvPath, *port, len(authTokens))
|
||||
} else {
|
||||
log.Printf("Starting Dynagate LLM Gateway. Config: %s, Port: %d, Auth: None (empty key file)", *csvPath, *port)
|
||||
}
|
||||
} else {
|
||||
log.Printf("Starting Dynagate LLM Gateway. Config: %s, Port: %d, Auth: None", *csvPath, *port)
|
||||
}
|
||||
@@ -62,9 +73,9 @@ func main() {
|
||||
|
||||
// Register routes
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/v1/models", handleModels(cm, authToken))
|
||||
mux.HandleFunc("/v1/chat/completions", handleChatCompletions(cm, authToken))
|
||||
mux.HandleFunc("/v1/images/generations", handleImageGenerations(cm, authToken))
|
||||
mux.HandleFunc("/v1/models", handleModels(cm, authTokens))
|
||||
mux.HandleFunc("/v1/chat/completions", handleChatCompletions(cm, authTokens))
|
||||
mux.HandleFunc("/v1/images/generations", handleImageGenerations(cm, authTokens))
|
||||
|
||||
server := &http.Server{
|
||||
Addr: fmt.Sprintf(":%d", *port),
|
||||
|
||||
Reference in New Issue
Block a user