From f09f55f743966205527cd9e56a4ddd9e162370d0 Mon Sep 17 00:00:00 2001 From: caydyan Date: Sun, 14 Jun 2026 22:03:42 +0800 Subject: [PATCH] Support configurable LLM providers --- .env.example | 5 ++++- README.md | 8 ++++++- helpers/llm_provider_helpers.py | 27 ++++++++++++++++++++++++ helpers/openai_helpers.py | 29 ++++++++++++++++++++----- main.py | 15 +++++++------ tests/test_llm_provider_helpers.py | 34 ++++++++++++++++++++++++++++++ 6 files changed, 105 insertions(+), 13 deletions(-) create mode 100644 helpers/llm_provider_helpers.py create mode 100644 tests/test_llm_provider_helpers.py diff --git a/.env.example b/.env.example index ab3da74..6df5350 100644 --- a/.env.example +++ b/.env.example @@ -1,9 +1,12 @@ AZURE_OPENAI_API_KEY=YOURKEY AZURE_OPENAI_ENDPOINT=https://YOURENDPOINT.openai.azure.com AZURE_OPENAI_API_VERSION=2024-02-01 +LLM_PROVIDER=azure_openai +OPENAI_API_KEY= +OPENAI_BASE_URL= MODEL=gpt-4o-mini EMBEDDING=text-embedding-3-small EMBEDDING_MODEL_MAX_TOKENS=8192 GITHUB_TOKEN=YOURTOKEN SUPABASE_URL=https://YOURLINK.supabase.co -SUPABASE_KEY=your_supabase_key \ No newline at end of file +SUPABASE_KEY=your_supabase_key diff --git a/README.md b/README.md index 4325025..83da1e7 100644 --- a/README.md +++ b/README.md @@ -55,6 +55,9 @@ To run this project in Daytona, you'll need to have Daytona installed. Follow th AZURE_OPENAI_ENDPOINT=your_azure_openai_endpoint AZURE_OPENAI_API_KEY=your_azure_openai_api_key AZURE_OPENAI_API_VERSION=your_azure_openai_api_version + LLM_PROVIDER=azure_openai + OPENAI_API_KEY=your_openai_api_key_if_using_openai_provider + OPENAI_BASE_URL=optional_openai_compatible_base_url MODEL=your_model_name GITHUB_TOKEN=your_github_token ``` @@ -84,6 +87,9 @@ Ensure the following environment variables are set in your `.env` file: AZURE_OPENAI_ENDPOINT=your_azure_openai_endpoint AZURE_OPENAI_API_KEY=your_azure_openai_api_key AZURE_OPENAI_API_VERSION=your_azure_openai_api_version +LLM_PROVIDER=azure_openai +OPENAI_API_KEY=your_openai_api_key_if_using_openai_provider +OPENAI_BASE_URL=optional_openai_compatible_base_url MODEL=your_model_name GITHUB_TOKEN=your_github_token ``` @@ -128,4 +134,4 @@ This project is licensed under the MIT License. See the [LICENSE](LICENSE) file --- -Thank you for using **devcontainer-generator**! If you have any questions or issues, feel free to open an issue on GitHub. \ No newline at end of file +Thank you for using **devcontainer-generator**! If you have any questions or issues, feel free to open an issue on GitHub. diff --git a/helpers/llm_provider_helpers.py b/helpers/llm_provider_helpers.py new file mode 100644 index 0000000..56863cd --- /dev/null +++ b/helpers/llm_provider_helpers.py @@ -0,0 +1,27 @@ +import os + + +DEFAULT_LLM_PROVIDER = "azure_openai" +SUPPORTED_LLM_PROVIDERS = {"azure_openai", "openai"} + + +def get_llm_provider(provider=None): + selected_provider = provider or os.getenv("LLM_PROVIDER", DEFAULT_LLM_PROVIDER) + return selected_provider.strip().lower().replace("-", "_") + + +def is_supported_llm_provider(provider=None): + return get_llm_provider(provider) in SUPPORTED_LLM_PROVIDERS + + +def required_env_vars_for_provider(provider=None): + selected_provider = get_llm_provider(provider) + if selected_provider == "azure_openai": + return [ + "AZURE_OPENAI_ENDPOINT", + "AZURE_OPENAI_API_KEY", + "AZURE_OPENAI_API_VERSION", + ] + if selected_provider == "openai": + return ["OPENAI_API_KEY"] + return ["LLM_PROVIDER"] diff --git a/helpers/openai_helpers.py b/helpers/openai_helpers.py index 4ba755a..24ea660 100644 --- a/helpers/openai_helpers.py +++ b/helpers/openai_helpers.py @@ -1,8 +1,10 @@ import os import logging -from openai import AzureOpenAI +from openai import AzureOpenAI, OpenAI import instructor +from helpers.llm_provider_helpers import get_llm_provider, is_supported_llm_provider, required_env_vars_for_provider + def setup_azure_openai(): logging.info("Setting up Azure OpenAI client...") return AzureOpenAI( @@ -11,18 +13,35 @@ def setup_azure_openai(): api_version=os.getenv("AZURE_OPENAI_API_VERSION"), ) +def setup_openai(): + logging.info("Setting up OpenAI-compatible client...") + return OpenAI( + api_key=os.getenv("OPENAI_API_KEY"), + base_url=os.getenv("OPENAI_BASE_URL") or None, + ) + +def setup_llm_client(): + provider = get_llm_provider() + if provider == "azure_openai": + return setup_azure_openai() + if provider == "openai": + return setup_openai() + raise ValueError(f"Unsupported LLM_PROVIDER: {provider}") + def setup_instructor(openai_client): logging.info("Setting up Instructor client...") return instructor.patch(openai_client) def check_env_vars(): + if not is_supported_llm_provider(): + print(f"Unsupported LLM_PROVIDER: {get_llm_provider()}") + return False + required_vars = [ - "AZURE_OPENAI_ENDPOINT", - "AZURE_OPENAI_API_KEY", - "AZURE_OPENAI_API_VERSION", "MODEL", "GITHUB_TOKEN", ] + required_vars.extend(required_env_vars_for_provider()) missing_vars = [var for var in required_vars if not os.environ.get(var)] if missing_vars: print( @@ -30,4 +49,4 @@ def check_env_vars(): "Please configure the env vars file properly." ) return False - return True \ No newline at end of file + return True diff --git a/main.py b/main.py index 2e6410d..8cb14b3 100644 --- a/main.py +++ b/main.py @@ -6,9 +6,10 @@ from dotenv import load_dotenv from supabase_client import supabase -from helpers.openai_helpers import setup_azure_openai, setup_instructor +from helpers.openai_helpers import setup_llm_client, setup_instructor from helpers.github_helpers import fetch_repo_context, check_url_exists from helpers.devcontainer_helpers import generate_devcontainer_json, validate_devcontainer_json +from helpers.llm_provider_helpers import get_llm_provider, is_supported_llm_provider, required_env_vars_for_provider from helpers.token_helpers import count_tokens, truncate_to_token_limit from models import DevContainer from schemas import DevContainerModel @@ -21,15 +22,17 @@ load_dotenv() def check_env_vars(): + if not is_supported_llm_provider(): + print(f"Unsupported LLM_PROVIDER: {get_llm_provider()}") + return False + required_vars = [ - "AZURE_OPENAI_ENDPOINT", - "AZURE_OPENAI_API_KEY", - "AZURE_OPENAI_API_VERSION", "MODEL", "GITHUB_TOKEN", "SUPABASE_URL", "SUPABASE_KEY", ] + required_vars.extend(required_env_vars_for_provider()) missing_vars = [var for var in required_vars if not os.environ.get(var)] if missing_vars: print(f"Missing environment variables: {', '.join(missing_vars)}. Please configure the env vars file properly.") @@ -202,9 +205,9 @@ async def get(fname:str, ext:str): # Initialize clients if check_env_vars(): - openai_client = setup_azure_openai() + openai_client = setup_llm_client() instructor_client = setup_instructor(openai_client) if __name__ == "__main__": logging.info("Starting FastHTML app...") - serve() \ No newline at end of file + serve() diff --git a/tests/test_llm_provider_helpers.py b/tests/test_llm_provider_helpers.py new file mode 100644 index 0000000..6226cd2 --- /dev/null +++ b/tests/test_llm_provider_helpers.py @@ -0,0 +1,34 @@ +import unittest + +from helpers.llm_provider_helpers import get_llm_provider, is_supported_llm_provider, required_env_vars_for_provider + + +class LlmProviderHelpersTest(unittest.TestCase): + def test_provider_names_are_normalized(self): + self.assertEqual(get_llm_provider("azure-openai"), "azure_openai") + self.assertEqual(get_llm_provider(" OpenAI "), "openai") + + def test_azure_openai_requires_azure_environment(self): + self.assertEqual( + required_env_vars_for_provider("azure_openai"), + [ + "AZURE_OPENAI_ENDPOINT", + "AZURE_OPENAI_API_KEY", + "AZURE_OPENAI_API_VERSION", + ], + ) + + def test_openai_compatible_provider_requires_openai_key(self): + self.assertEqual(required_env_vars_for_provider("openai"), ["OPENAI_API_KEY"]) + + def test_unknown_provider_reports_llm_provider_as_missing(self): + self.assertEqual(required_env_vars_for_provider("anthropic"), ["LLM_PROVIDER"]) + + def test_supported_provider_check(self): + self.assertTrue(is_supported_llm_provider("azure_openai")) + self.assertTrue(is_supported_llm_provider("openai")) + self.assertFalse(is_supported_llm_provider("anthropic")) + + +if __name__ == "__main__": + unittest.main()