diff --git a/config.toml b/config.toml index 89e94ed52..d3cab7969 100644 --- a/config.toml +++ b/config.toml @@ -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" diff --git a/src/rai_core/pyproject.toml b/src/rai_core/pyproject.toml index d7c3e5ae8..74e3afd7b 100644 --- a/src/rai_core/pyproject.toml +++ b/src/rai_core/pyproject.toml @@ -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", diff --git a/src/rai_core/rai/initialization/config_initialization.py b/src/rai_core/rai/initialization/config_initialization.py index 94028b252..525ec110b 100644 --- a/src/rai_core/rai/initialization/config_initialization.py +++ b/src/rai_core/rai/initialization/config_initialization.py @@ -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" diff --git a/src/rai_core/rai/initialization/model_initialization.py b/src/rai_core/rai/initialization/model_initialization.py index 6cd5d7e42..6fbd4ee39 100644 --- a/src/rai_core/rai/initialization/model_initialization.py +++ b/src/rai_core/rai/initialization/model_initialization.py @@ -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 @@ -67,6 +69,11 @@ class GoogleConfig(ModelConfig): pass +@dataclass +class MistralAIConfig(ModelConfig): + base_url: str + + @dataclass class LangfuseConfig: use_langfuse: bool @@ -93,6 +100,7 @@ class RAIConfig: openai: OpenAIConfig ollama: OllamaConfig google: GoogleConfig + mistralai: MistralAIConfig tracing: TracingConfig @@ -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=""), @@ -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( @@ -156,6 +172,7 @@ def load_config(config_path: Optional[str] = None) -> RAIConfig: openai=openai, ollama=ollama, google=google, + mistralai=mistralai, tracing=tracing, ) @@ -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 ) @@ -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}") @@ -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) @@ -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}") @@ -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}") diff --git a/tests/conftest.py b/tests/conftest.py index 7e81124cc..30a149cdf 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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" @@ -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) diff --git a/tests/initialization/test_model_initialization.py b/tests/initialization/test_model_initialization.py index eafdad88e..2e94683a9 100644 --- a/tests/initialization/test_model_initialization.py +++ b/tests/initialization/test_model_initialization.py @@ -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"