How to Integrate Custom LLM Providers in OpenRisk: A Complete Guide
You can integrate custom LLM providers into OpenRisk by implementing the LLMProvider abstract base class and registering it with the ProviderRegistry decorator, enabling automatic discovery by the LLMClient.
OpenRisk (derisk-ai/openderisk) uses a plug-in architecture that makes it straightforward to integrate custom LLM providers beyond the default OpenAI configuration. Whether you need to connect to a proprietary API, a self-hosted model, or an emerging service, the provider system abstracts away the complexity through standardized interfaces.
Understanding the Provider Architecture
The integration relies on three core components working together to resolve and execute LLM requests.
LLMProvider Abstract Base Class
Located at derisk/agent/util/llm/provider/base.py, the LLMProvider abstract base class defines the contract that every provider must implement. It requires four asynchronous methods:
generate– Executes a single-shot completion request and returns aModelOutput.generate_stream– YieldsModelOutputchunks for streaming responses.models– Returns a list of availableModelMetadataobjects.count_token– Estimates token counts for the given model and prompt.
ProviderRegistry Global Singleton
The ProviderRegistry in derisk/agent/util/llm/provider/provider_registry.py acts as a factory and lookup table. It maps provider names (lowercase strings) to concrete classes and stores the environment variable key for API credentials.
When you decorate a class with @ProviderRegistry.register("name", env_key="ENV_VAR"), the registry stores the mapping and retrieves the API key from the environment if not explicitly provided.
LLMClient Consumer
The LLMClient in derisk/agent/util/llm/llm_client.py resolves the correct provider at runtime. It extracts the provider field from a ModelRequest, queries the ProviderRegistry for the corresponding class, and instantiates it with the appropriate api_key and configuration parameters.
Step-by-Step: Creating a Custom LLM Provider
Follow these steps to integrate a new LLM backend into the OpenRisk ecosystem.
Step 1: Implement the LLMProvider Abstract Base Class
Create a new module in packages/derisk-core/src/derisk/agent/util/llm/provider/. Inherit from LLMProvider and implement all required abstract methods.
from typing import AsyncIterator, List
from derisk.core.interface.llm import ModelRequest, ModelOutput, ModelMetadata
from derisk.agent.util.llm.provider.base import LLMProvider
class MyCustomProvider(LLMProvider):
async def generate(self, request: ModelRequest) -> ModelOutput:
# Implement single-shot generation
pass
async def generate_stream(self, request: ModelRequest) -> AsyncIterator[ModelOutput]:
# Implement streaming or raise NotImplementedError
pass
async def models(self) -> List[ModelMetadata]:
# Return available models
pass
async def count_token(self, model: str, prompt: str) -> int:
# Return token count estimate
pass
Step 2: Register Your Provider with the ProviderRegistry
Use the @ProviderRegistry.register decorator to make your provider discoverable. Specify the provider name and the environment variable key for the API key.
from derisk.agent.util.llm.provider.provider_registry import ProviderRegistry
@ProviderRegistry.register("mycustom", env_key="MYCUSTOM_API_KEY")
class MyCustomProvider(LLMProvider):
def __init__(self, api_key: str, base_url: str = None, **kwargs):
self.api_key = api_key
self.base_url = base_url or "https://api.mycustom.ai/v1"
# Initialize HTTP client or SDK here
Step 3: Handle Authentication and Configuration
The LLMClient passes api_key, base_url, model, and any extra kwargs from ModelRequest.extra_kwargs to your provider's constructor. If api_key is not provided, the registry reads it from the env_key you specified.
# The registry instantiation logic roughly follows this pattern:
provider_class = ProviderRegistry.get_provider_class("mycustom")
api_key = api_key or os.getenv(ProviderRegistry.get_env_key("mycustom"))
provider = provider_class(api_key=api_key, base_url=base_url, model=model, **extra_kwargs)
Complete Implementation Example
Here is a minimal, runnable implementation of a custom HTTP-based provider that follows the OpenRisk conventions.
# packages/derisk-core/src/derisk/agent/util/llm/provider/myprovider.py
import httpx
import json
import logging
from typing import AsyncIterator, List
from derisk.core.interface.llm import ModelRequest, ModelOutput, ModelMetadata
from derisk.agent.util.llm.provider.base import LLMProvider
from derisk.agent.util.llm.provider.provider_registry import ProviderRegistry
logger = logging.getLogger(__name__)
@ProviderRegistry.register("myprovider", env_key="MY_API_KEY")
class MyProvider(LLMProvider):
"""Example HTTP-based LLM provider for OpenRisk integration."""
def __init__(self, api_key: str, base_url: str = "https://api.myai.com/v1", **_: dict):
self.api_key = api_key
self.base_url = base_url
self.client = httpx.AsyncClient(timeout=30)
async def _post(self, endpoint: str, payload: dict) -> dict:
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json"
}
response = await self.client.post(
f"{self.base_url}/{endpoint}",
headers=headers,
json=payload
)
response.raise_for_status()
return response.json()
async def generate(self, request: ModelRequest) -> ModelOutput:
payload = {
"model": request.model,
"messages": request.to_common_messages(support_system_role=True),
"temperature": request.temperature,
"max_tokens": request.max_new_tokens,
}
if request.tools:
payload["tools"] = request.tools
data = await self._post("chat/completions", payload)
choice = data["choices"][0]
return ModelOutput(
error_code=0,
text=choice["message"]["content"],
tool_calls=choice["message"].get("tool_calls"),
finish_reason=choice.get("finish_reason"),
usage=data.get("usage"),
)
async def generate_stream(self, request: ModelRequest) -> AsyncIterator[ModelOutput]:
# Fallback to non-streaming if the backend does not support SSE
result = await self.generate(request)
yield result
async def models(self) -> List[ModelMetadata]:
data = await self._post("models", {})
return [ModelMetadata(model=m["id"]) for m in data.get("data", [])]
async def count_token(self, model: str, prompt: str) -> int:
# Rough heuristic; replace with provider-specific logic if available
return len(prompt) // 4
Advanced Integration Patterns
Using Factory Functions for Multi-Tenant Setups
If your provider requires complex initialization logic—such as rotating API keys based on tenant ID or connection pooling—use the factory parameter in the registry decorator.
def my_factory(**kwargs):
# Implement custom logic to select API keys
tenant_id = kwargs.get("tenant_id")
key = select_key_for_tenant(tenant_id)
return MyProvider(api_key=key, **kwargs)
@ProviderRegistry.register("myprovider", factory=my_factory, env_key="MY_API_KEY")
class MyProvider(LLMProvider):
...
The ProviderRegistry.create_provider method will invoke my_factory instead of the class constructor, passing all configuration parameters.
Handling Non-Streaming Providers
Not all LLM APIs support Server-Sent Events (SSE) or chunked streaming. For these cases, implement generate_stream as a fallback that yields the single result from generate.
async def generate_stream(self, request: ModelRequest) -> AsyncIterator[ModelOutput]:
# Fallback for APIs without streaming support
result = await self.generate(request)
yield result
This ensures compatibility with the LLMClient streaming interface while supporting providers that only offer synchronous endpoints.
Testing Your Custom Provider
Before deploying, validate your implementation using the test patterns established in the OpenRisk codebase.
# tests/test_myprovider.py
import pytest
from derisk.agent.util.llm.llm_client import LLMClient
from derisk.agent.util.llm.provider.provider_registry import ProviderRegistry
from derisk.core.interface.llm import ModelRequest, ModelOutput
@pytest.mark.asyncio
async def test_custom_provider_registration():
client = LLMClient()
# Verify the provider is discoverable
assert ProviderRegistry.get_provider_class("myprovider") is not None
# Test basic generation
req = ModelRequest(
provider="myprovider",
model="test-model",
messages=[{"role": "user", "content": "Hello"}],
api_key="test-key" # Or set MY_API_KEY env var
)
response = await client.request(req)
assert isinstance(response, ModelOutput)
assert response.error_code == 0
Summary
Integrating custom LLM providers into OpenRisk requires implementing four key steps:
- Implement the
LLMProviderinterface inderisk/agent/util/llm/provider/base.pyby overridinggenerate,generate_stream,models, andcount_token. - Register with
ProviderRegistryusing the@ProviderRegistry.registerdecorator to map a provider name to your class and specify the environment variable for API keys. - Handle authentication through the constructor signature
(api_key, base_url, **kwargs), allowingLLMClientto pass credentials fromModelRequestor environment variables. - Test integration by instantiating
LLMClientand sending aModelRequestwith your registered provider name to verify end-to-end functionality.
Frequently Asked Questions
How do I override the default OpenAI provider with my custom implementation?
Set the provider field in your ModelRequest to your registered provider name. The LLMClient in derisk/agent/util/llm/llm_client.py resolves the provider from the ProviderRegistry based on this field. If you want your provider to be the system default, you would need to modify the configuration layer that instantiates ModelRequest objects to use your provider name instead of "openai".
Can I integrate a provider that uses a proprietary SDK instead of HTTP requests?
Yes. The LLMProvider interface is transport-agnostic. In your implementation file, import the vendor's SDK and use it within the generate and generate_stream methods. Initialize the SDK client in your __init__ method using the api_key and base_url parameters passed from the registry. As long as you return ModelOutput objects with the correct fields, the rest of the OpenRisk agent system will function normally.
What happens if my LLM service does not support token counting?
Implement the count_token method with a heuristic fallback. The analysis shows that len(prompt) // 4 serves as a reasonable approximation for many models. Return this estimate in your implementation. While not perfectly accurate, this allows the agent system to enforce context window limits and budget constraints without requiring provider-specific tokenizers. If your service later adds a token counting endpoint, you can update the method to call that API instead.
How do I handle different API keys for different users in a multi-tenant deployment?
Use the factory parameter in the @ProviderRegistry.register decorator. Define a factory function that accepts **kwargs, extracts a tenant_id or user identifier from the keyword arguments, and selects the appropriate API key from your key management system. The factory then instantiates your provider class with the selected key. This pattern is demonstrated in the Theta provider implementation, where complex initialization logic and quota handling are encapsulated in a factory function.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →