Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions config.toml
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,12 @@ simple_model = "gemini-3-flash"
complex_model = "gemini-3-pro"
embeddings_model = "text-embedding-004"

[mistralai]
simple_model = "mistral-small-latest"
complex_model = "mistral-large-latest"
embeddings_model = "mistral-embed"
base_url = "https://api.mistral.ai/v1/"

[tracing]
project = "rai"

Expand Down
1 change: 1 addition & 0 deletions src/rai_core/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ dependencies = [
"langchain-openai",
"langchain-ollama",
"langchain-google-genai",
"langchain-mistralai",
"langchain-community",
"requests>=2.32.2,<3.0.0",
"coloredlogs>=15.0.1,<16.0.0",
Expand Down
6 changes: 6 additions & 0 deletions src/rai_core/rai/initialization/config_initialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,12 @@
complex_model = "gemini-3-pro"
embeddings_model = "text-embedding-004"

[mistralai]
simple_model = "mistral-small-latest"
complex_model = "mistral-large-latest"
embeddings_model = "mistral-embed"
base_url = "https://api.mistral.ai/v1/"

[tracing]
project = "rai"

Expand Down
52 changes: 50 additions & 2 deletions src/rai_core/rai/initialization/model_initialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
from langchain_core.callbacks.base import BaseCallbackHandler
from langchain_core.embeddings import Embeddings
from langchain_core.tracers.langchain import LangChainTracer
from langchain_google_genai import ChatGoogleGenerativeAI
from langchain_mistralai import ChatMistralAI
from langchain_ollama import ChatOllama
from langchain_openai import ChatOpenAI
from langsmith import Client
Expand Down Expand Up @@ -67,6 +69,11 @@ class GoogleConfig(ModelConfig):
pass


@dataclass
class MistralAIConfig(ModelConfig):
base_url: str


@dataclass
class LangfuseConfig:
use_langfuse: bool
Expand All @@ -93,6 +100,7 @@ class RAIConfig:
openai: OpenAIConfig
ollama: OllamaConfig
google: GoogleConfig
mistralai: MistralAIConfig
tracing: TracingConfig


Expand All @@ -107,6 +115,9 @@ class RAIConfig:
simple_model="", complex_model="", embeddings_model="", base_url=""
)
_DEFAULT_GOOGLE = GoogleConfig(simple_model="", complex_model="", embeddings_model="")
_DEFAULT_MISTRALAI = MistralAIConfig(
simple_model="", complex_model="", embeddings_model="", base_url=""
)
_DEFAULT_TRACING = TracingConfig(
project="",
langfuse=LangfuseConfig(use_langfuse=False, host=""),
Expand Down Expand Up @@ -140,6 +151,11 @@ def load_config(config_path: Optional[str] = None) -> RAIConfig:
if "google" in config_dict
else _DEFAULT_GOOGLE
)
mistralai = (
MistralAIConfig(**config_dict["mistralai"])
if "mistralai" in config_dict
else _DEFAULT_MISTRALAI
)

if "tracing" in config_dict:
tracing = TracingConfig(
Expand All @@ -156,6 +172,7 @@ def load_config(config_path: Optional[str] = None) -> RAIConfig:
openai=openai,
ollama=ollama,
google=google,
mistralai=mistralai,
tracing=tracing,
)

Expand All @@ -181,7 +198,7 @@ def get_llm_model(
vendor: Optional[str] = None,
config_path: Optional[str] = None,
**kwargs: Any,
) -> ChatOpenAI | ChatBedrock | ChatOllama | Any:
) -> ChatOpenAI | ChatBedrock | ChatOllama | ChatGoogleGenerativeAI | ChatMistralAI:
model_config, vendor = get_llm_model_config_and_vendor(
model_type, vendor, config_path
)
Expand Down Expand Up @@ -213,6 +230,12 @@ def get_llm_model(

model_config = cast(GoogleConfig, model_config)
return ChatGoogleGenerativeAI(model=model, **kwargs)
elif vendor == "mistralai":
from langchain_mistralai import ChatMistralAI

model_config = cast(MistralAIConfig, model_config)

return ChatMistralAI(model=model, base_url=model_config.base_url, **kwargs)
else:
raise ValueError(f"Unknown LLM vendor: {vendor}")

Expand All @@ -222,7 +245,7 @@ def get_llm_model_direct(
vendor: str,
config_path: Optional[str] = None,
**kwargs: Any,
) -> ChatOpenAI | ChatBedrock | ChatOllama | Any:
) -> ChatOpenAI | ChatBedrock | ChatOllama | ChatGoogleGenerativeAI | ChatMistralAI:
config = load_config(config_path)
model_config = getattr(config, vendor)

Expand Down Expand Up @@ -255,6 +278,12 @@ def get_llm_model_direct(

model_config = cast(GoogleConfig, model_config)
return ChatGoogleGenerativeAI(model=model_name, **kwargs)
elif vendor == "mistralai":
from langchain_mistralai import ChatMistralAI

model_config = cast(MistralAIConfig, model_config)

return ChatMistralAI(model=model, base_url=model_config.base_url, **kwargs)
else:
raise ValueError(f"Unknown LLM vendor: {vendor}")

Expand Down Expand Up @@ -352,6 +381,25 @@ def get_embeddings_model(
"vendor": vendor,
}
return embeddings
elif vendor == "mistralai":
from langchain_mistralai import MistralAIEmbeddings

model_config = cast(MistralAIConfig, model_config)
embeddings = MistralAIEmbeddings(model=model_config.embeddings_model)
if return_kwargs:
c = (
str(embeddings.__class__)
.strip("<>")
.replace("class '", "")
.replace("'", "")
)
return embeddings, {
"class": c,
"model": model_config.embeddings_model,
"base_url": model_config.base_url,
"vendor": vendor,
}
return embeddings
else:
raise ValueError(f"Unknown embeddings vendor: {vendor}")

Expand Down
12 changes: 12 additions & 0 deletions tests/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,12 @@ def _create_config(langfuse_enabled=False, langsmith_enabled=False):
complex_model = "gemini-3-pro"
embeddings_model = "text-embedding-004"

[mistralai]
simple_model = "mistral-small-latest"
complex_model = "mistral-large-latest"
embeddings_model = "mistral-embed"
base_url = "https://api.mistral.ai/v1/"

[tracing]
project = "test-project"

Expand Down Expand Up @@ -147,6 +153,12 @@ def _create_config():
simple_model = "gemini-3-flash"
complex_model = "gemini-3-pro"
embeddings_model = "text-embedding-004"

[mistralai]
simple_model = "mistral-small-latest"
complex_model = "mistral-large-latest"
embeddings_model = "mistral-embed"
base_url = "https://api.mistral.ai/v1/"
"""

f.write(config_content)
Expand Down
6 changes: 6 additions & 0 deletions tests/initialization/test_model_initialization.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,12 @@
complex_model = "gemini-3-pro"
embeddings_model = "text-embedding-004"

[mistralai]
simple_model = "mistral-small-latest"
complex_model = "mistral-large-latest"
embeddings_model = "mistral-embed"
base_url = "https://api.mistral.ai/v1/"

[tracing]
project = "rai"

Expand Down
Loading