Merge pull request #43 from hezhizhen/custom-endpoint

feat: support custom endpoint with openai compatible apis
This commit is contained in:
rahul 2025-08-19 12:43:06 -04:00 committed by GitHub
commit 25e98daef3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 23 additions and 7 deletions

View file

@ -7,6 +7,7 @@ GEMINI_API_KEY="your-google-ai-studio-key"
# Controls which provider (google or openai) is preferred for mapping haiku/sonnet. # Controls which provider (google or openai) is preferred for mapping haiku/sonnet.
# Defaults to openai if not set. # Defaults to openai if not set.
PREFERRED_PROVIDER="openai" PREFERRED_PROVIDER="openai"
OPENAI_BASE_URL="https://api.openai.com/v1"
# Optional: Specify the exact models to map haiku/sonnet to. # Optional: Specify the exact models to map haiku/sonnet to.
# If PREFERRED_PROVIDER=google, these MUST be valid Gemini model names known to the server. # If PREFERRED_PROVIDER=google, these MUST be valid Gemini model names known to the server.

View file

@ -82,6 +82,9 @@ ANTHROPIC_API_KEY = os.environ.get("ANTHROPIC_API_KEY")
OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY") OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY")
GEMINI_API_KEY = os.environ.get("GEMINI_API_KEY") GEMINI_API_KEY = os.environ.get("GEMINI_API_KEY")
# Get OpenAI base URL from environment (if set)
OPENAI_BASE_URL = os.environ.get("OPENAI_BASE_URL")
# Get preferred provider (default to openai) # Get preferred provider (default to openai)
PREFERRED_PROVIDER = os.environ.get("PREFERRED_PROVIDER", "openai").lower() PREFERRED_PROVIDER = os.environ.get("PREFERRED_PROVIDER", "openai").lower()
@ -1106,7 +1109,12 @@ async def create_message(
# Determine which API key to use based on the model # Determine which API key to use based on the model
if request.model.startswith("openai/"): if request.model.startswith("openai/"):
litellm_request["api_key"] = OPENAI_API_KEY litellm_request["api_key"] = OPENAI_API_KEY
logger.debug(f"Using OpenAI API key for model: {request.model}") # Use custom OpenAI base URL if configured
if OPENAI_BASE_URL:
litellm_request["api_base"] = OPENAI_BASE_URL
logger.debug(f"Using OpenAI API key and custom base URL {OPENAI_BASE_URL} for model: {request.model}")
else:
logger.debug(f"Using OpenAI API key for model: {request.model}")
elif request.model.startswith("gemini/"): elif request.model.startswith("gemini/"):
litellm_request["api_key"] = GEMINI_API_KEY litellm_request["api_key"] = GEMINI_API_KEY
logger.debug(f"Using Gemini API key for model: {request.model}") logger.debug(f"Using Gemini API key for model: {request.model}")
@ -1386,11 +1394,18 @@ async def count_tokens(
200 # Assuming success at this point 200 # Assuming success at this point
) )
# Prepare token counter arguments
token_counter_args = {
"model": converted_request["model"],
"messages": converted_request["messages"],
}
# Add custom base URL for OpenAI models if configured
if request.model.startswith("openai/") and OPENAI_BASE_URL:
token_counter_args["api_base"] = OPENAI_BASE_URL
# Count tokens # Count tokens
token_count = token_counter( token_count = token_counter(**token_counter_args)
model=converted_request["model"],
messages=converted_request["messages"],
)
# Return Anthropic-style response # Return Anthropic-style response
return TokenCountResponse(input_tokens=token_count) return TokenCountResponse(input_tokens=token_count)