chATLAS_Chains 0.1.6__tar.gz → 0.1.7__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (48) hide show
  1. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/PKG-INFO +72 -7
  2. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/README.md +71 -6
  3. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/advanced.py +28 -29
  4. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/basic.py +9 -18
  5. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/basic_graph.py +2 -6
  6. chatlas_chains-0.1.7/chATLAS_Chains/chains/conversational_graph.py +383 -0
  7. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/documents/rerank.py +2 -1
  8. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/llm/groq.py +1 -1
  9. chatlas_chains-0.1.7/chATLAS_Chains/llm/model_selection.py +302 -0
  10. chatlas_chains-0.1.7/chATLAS_Chains/llm/runnables.py +156 -0
  11. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/starters.py +136 -0
  12. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/search/basic.py +1 -1
  13. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/vectorstore.py +4 -3
  14. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/PKG-INFO +72 -7
  15. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/SOURCES.txt +6 -1
  16. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/requires.txt +1 -1
  17. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/pyproject.toml +2 -2
  18. chatlas_chains-0.1.7/tests/conftest.py +99 -0
  19. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/tests/test_chains.py +11 -15
  20. chatlas_chains-0.1.7/tests/test_chat_model_kwargs.py +168 -0
  21. chatlas_chains-0.1.7/tests/test_conversational.py +238 -0
  22. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/tests/test_groq.py +0 -1
  23. chatlas_chains-0.1.7/tests/test_llm_runnables.py +48 -0
  24. chatlas_chains-0.1.7/tests/test_model_selection.py +129 -0
  25. chatlas_chains-0.1.6/chATLAS_Chains/llm/model_selection.py +0 -176
  26. chatlas_chains-0.1.6/tests/conftest.py +0 -68
  27. chatlas_chains-0.1.6/tests/test_llm.py +0 -227
  28. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/LICENSE +0 -0
  29. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/__init__.py +0 -0
  30. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/__init__.py +0 -0
  31. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
  32. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
  33. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/documents/rrf.py +0 -0
  34. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/llm/__init__.py +0 -0
  35. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/log.py +0 -0
  36. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/__init__.py +0 -0
  37. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  38. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/query/query_rewriting.py +0 -0
  39. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/search/__init__.py +0 -0
  40. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/utils/__init__.py +0 -0
  41. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains/utils/doc_utils.py +0 -0
  42. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  43. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/top_level.txt +0 -0
  44. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/setup.cfg +0 -0
  45. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/tests/__init__.py +0 -0
  46. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/tests/test_rrf.py +0 -0
  47. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/tests/test_search.py +0 -0
  48. {chatlas_chains-0.1.6 → chatlas_chains-0.1.7}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: chATLAS_Chains
3
- Version: 0.1.6
3
+ Version: 0.1.7
4
4
  Summary: A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project.
5
5
  Author-email: Joe Egan <joseph.caimin.egan@cern.ch>
6
6
  License: Apache-2.0
@@ -19,9 +19,9 @@ Requires-Dist: langgraph
19
19
  Requires-Dist: sentence-transformers>=3.0.0
20
20
  Requires-Dist: tiktoken
21
21
  Requires-Dist: pinecone
22
- Requires-Dist: torch==2.2.1
23
22
  Requires-Dist: chatlas-embed
24
23
  Requires-Dist: psycopg2-binary>=2.9.10
24
+ Requires-Dist: httpx[socks]>=0.28.1
25
25
  Dynamic: license-file
26
26
 
27
27
 
@@ -62,14 +62,71 @@ export CHATLAS_GROQ_KEY="your groq api key"
62
62
 
63
63
  **note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
64
64
  ```bash
65
- ssh -L 3000:cs-513-ml003:3000 <LXPLUS_USERNAME>@lxplus.cern.ch
65
+ ssh -L 3000:cs-513-ml003:3000 $LXPLUS_USERNAME@lxplus.cern.ch
66
66
  export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
67
67
  ```
68
68
 
69
- ## Available Chains
70
- - chains.basic.basic_retrieval_chain
71
- - chains.basic_graph.basic_retrieval_graph
72
- - chains.advanced.advanced_rag
69
+ 3. Using LLMs via CERN's LiteLLM API, here is the [repo](https://gitlab.cern.ch/itgpt/litellm-okd/-/tree/main) and some [setup instructions](https://codimd.web.cern.ch/tQKiMa13Q4O-EJXWTO3N7w?view#Using-Your-Dedicated-API-Key-to-Access-LLMs) for reference.
70
+ ```bash
71
+ export CHATLAS_CHAINS_LITELLM_KEY="your litellm key"
72
+ ```
73
+
74
+ ## Supported Chains
75
+
76
+ More details [here](chATLAS_Chains/chains/README.md)
77
+
78
+ - `chains.basic.basic_retrieval_chain`
79
+ - `chains.advanced.advanced_rag`
80
+
81
+ ### Model Configuration in Chains
82
+
83
+ Supported chain constructors now accept a typed `chat_model_kwargs` argument for model options (for example:
84
+ `temperature`, `max_tokens`, `service_provider`, `api_key`, `base_url`, `proxy`).
85
+
86
+ ```python
87
+ from chATLAS_Chains.chains.basic import basic_retrieval_chain
88
+
89
+ chain = basic_retrieval_chain(
90
+ prompt=...,
91
+ vectorstore=...,
92
+ model_name="gpt-4o-mini",
93
+ chat_model_kwargs={"temperature": 0.1, "max_tokens": 512},
94
+ )
95
+ ```
96
+
97
+ ## Forwarding vectorstore connections
98
+
99
+ If not on the CERN network, you can forward the connection to the postgres servers with:
100
+
101
+ ```bash
102
+ ssh -N \
103
+ -L 6624:dbod-chatlas.cern.ch:6624 \
104
+ -L 6606:dbod-chatlas-cds.cern.ch:6606 \
105
+ "$LXPLUS_USERNAME"@lxplus.cern.ch
106
+ export CHATLAS_PORT_FORWARDING=1
107
+ ```
108
+
109
+ You can then the helper function [`get_vectorstore`](chATLAS_Chains/vectorstore.py)
110
+
111
+ ## Testing Environment Variables
112
+
113
+ Some tests are DB-backed integration tests (`tests/test_chains.py`, `tests/test_conversational.py`, `tests/test_search.py`).
114
+ If the DB/test environment is not configured, these tests are skipped by `tests/conftest.py`.
115
+
116
+ `tests/conftest.py` now uses explicit controls:
117
+
118
+ - `CHATLAS_PORT_FORWARDING`: enable localhost DB tunnels (`1`, `true`, `True`)
119
+ - `CHATLAS_DB_PASSWORD`
120
+
121
+ ### Local Example (with DB tunnels)
122
+
123
+ ```bash
124
+ export CHATLAS_DB_PASSWORD="..."
125
+ export CHATLAS_PORT_FORWARDING=1
126
+ unset GITLAB_PAT
127
+
128
+ uv run pytest -q
129
+ ```
73
130
 
74
131
  ## Postgres
75
132
 
@@ -91,6 +148,14 @@ CREATE EXTENSION IF NOT EXISTS vector;
91
148
  ```
92
149
  ## CHANGELOG
93
150
 
151
+ #### 0.1.7
152
+
153
+ Support for CERN-hosted LiteLLM models
154
+
155
+ Multi-turn conversational RAG with (local) conversation history
156
+
157
+ Bugfixes
158
+
94
159
  #### 0.1.6
95
160
 
96
161
  Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
@@ -36,14 +36,71 @@ export CHATLAS_GROQ_KEY="your groq api key"
36
36
 
37
37
  **note** The API address is local to the CERN network. If not at CERN, you can forward it like so:
38
38
  ```bash
39
- ssh -L 3000:cs-513-ml003:3000 <LXPLUS_USERNAME>@lxplus.cern.ch
39
+ ssh -L 3000:cs-513-ml003:3000 $LXPLUS_USERNAME@lxplus.cern.ch
40
40
  export CHATLAS_GROQ_BASE_URL="http://localhost:3000"
41
41
  ```
42
42
 
43
- ## Available Chains
44
- - chains.basic.basic_retrieval_chain
45
- - chains.basic_graph.basic_retrieval_graph
46
- - chains.advanced.advanced_rag
43
+ 3. Using LLMs via CERN's LiteLLM API, here is the [repo](https://gitlab.cern.ch/itgpt/litellm-okd/-/tree/main) and some [setup instructions](https://codimd.web.cern.ch/tQKiMa13Q4O-EJXWTO3N7w?view#Using-Your-Dedicated-API-Key-to-Access-LLMs) for reference.
44
+ ```bash
45
+ export CHATLAS_CHAINS_LITELLM_KEY="your litellm key"
46
+ ```
47
+
48
+ ## Supported Chains
49
+
50
+ More details [here](chATLAS_Chains/chains/README.md)
51
+
52
+ - `chains.basic.basic_retrieval_chain`
53
+ - `chains.advanced.advanced_rag`
54
+
55
+ ### Model Configuration in Chains
56
+
57
+ Supported chain constructors now accept a typed `chat_model_kwargs` argument for model options (for example:
58
+ `temperature`, `max_tokens`, `service_provider`, `api_key`, `base_url`, `proxy`).
59
+
60
+ ```python
61
+ from chATLAS_Chains.chains.basic import basic_retrieval_chain
62
+
63
+ chain = basic_retrieval_chain(
64
+ prompt=...,
65
+ vectorstore=...,
66
+ model_name="gpt-4o-mini",
67
+ chat_model_kwargs={"temperature": 0.1, "max_tokens": 512},
68
+ )
69
+ ```
70
+
71
+ ## Forwarding vectorstore connections
72
+
73
+ If not on the CERN network, you can forward the connection to the postgres servers with:
74
+
75
+ ```bash
76
+ ssh -N \
77
+ -L 6624:dbod-chatlas.cern.ch:6624 \
78
+ -L 6606:dbod-chatlas-cds.cern.ch:6606 \
79
+ "$LXPLUS_USERNAME"@lxplus.cern.ch
80
+ export CHATLAS_PORT_FORWARDING=1
81
+ ```
82
+
83
+ You can then the helper function [`get_vectorstore`](chATLAS_Chains/vectorstore.py)
84
+
85
+ ## Testing Environment Variables
86
+
87
+ Some tests are DB-backed integration tests (`tests/test_chains.py`, `tests/test_conversational.py`, `tests/test_search.py`).
88
+ If the DB/test environment is not configured, these tests are skipped by `tests/conftest.py`.
89
+
90
+ `tests/conftest.py` now uses explicit controls:
91
+
92
+ - `CHATLAS_PORT_FORWARDING`: enable localhost DB tunnels (`1`, `true`, `True`)
93
+ - `CHATLAS_DB_PASSWORD`
94
+
95
+ ### Local Example (with DB tunnels)
96
+
97
+ ```bash
98
+ export CHATLAS_DB_PASSWORD="..."
99
+ export CHATLAS_PORT_FORWARDING=1
100
+ unset GITLAB_PAT
101
+
102
+ uv run pytest -q
103
+ ```
47
104
 
48
105
  ## Postgres
49
106
 
@@ -65,6 +122,14 @@ CREATE EXTENSION IF NOT EXISTS vector;
65
122
  ```
66
123
  ## CHANGELOG
67
124
 
125
+ #### 0.1.7
126
+
127
+ Support for CERN-hosted LiteLLM models
128
+
129
+ Multi-turn conversational RAG with (local) conversation history
130
+
131
+ Bugfixes
132
+
68
133
  #### 0.1.6
69
134
 
70
135
  Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
@@ -119,4 +184,4 @@ chATLAS_Benchmark is released under Apache v2.0 license.
119
184
 
120
185
  *For questions and support, please [contact](mailto:joseph.caimin.egan@cern.ch)*
121
186
 
122
- </div>
187
+ </div>
@@ -23,13 +23,17 @@ from langgraph.graph.state import CompiledStateGraph
23
23
 
24
24
  from chATLAS_Chains.documents.rerank import rerank_documents
25
25
  from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
26
- from chATLAS_Chains.llm.groq import RetryConfig
27
- from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS, get_chat_model
26
+ from chATLAS_Chains.llm.model_selection import (
27
+ GROQ_PRODUCTION_MODELS,
28
+ ChatModelKwargs,
29
+ get_chat_model,
30
+ sanitize_chat_model_kwargs,
31
+ )
28
32
  from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
29
33
  from chATLAS_Chains.query.query_rewriting import rewrite_query
30
34
  from chATLAS_Chains.search.basic import search_runnable
31
35
  from chATLAS_Chains.utils.doc_utils import combine_documents
32
- from chATLAS_Embed.Base import VectorStore
36
+ from chATLAS_Embed.VectorStores import VectorStore
33
37
 
34
38
  logger = logging.getLogger(__name__)
35
39
 
@@ -40,27 +44,24 @@ class HybridGraphState(TypedDict, total=False):
40
44
  search_kwargs: dict
41
45
  docs: list[Document]
42
46
  answer: str
47
+ chat_history: str # Optional
43
48
 
44
49
 
45
50
  def advanced_rag(
46
51
  vectorstore: VectorStore | list[VectorStore],
47
52
  model_name: str,
48
53
  prompt: str | None = None,
49
- max_tokens: int | None = None,
50
- temperature: float = 0.1,
51
- use_preview_models: bool = False,
54
+ chat_model_kwargs: ChatModelKwargs | None = None,
52
55
  enable_query_rewriting: bool = False,
53
56
  enable_rrf: bool = False,
54
57
  enable_reranking: bool = False,
55
58
  # enable_self_evaluation: bool = False,
56
59
  query_rewriting_model: str = GROQ_PRODUCTION_MODELS[0],
57
- query_rewriting_temperature: float = 0.1,
60
+ query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
58
61
  rerank_model: str = "cohere-rerank-3.5",
59
62
  pinecone_api_key: str | None = None,
60
63
  rrf_constant: float = 60.0,
61
64
  rrf_weights: dict[str, float] | None = None,
62
- retry_config: RetryConfig | None = None,
63
- fallback_models: list[str] | None = None,
64
65
  ) -> CompiledStateGraph:
65
66
  """
66
67
  Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
@@ -68,20 +69,17 @@ def advanced_rag(
68
69
  :param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
69
70
  :param vectorstore: Single vectorstore instance or list of vectorstore instances to search
70
71
  :param model_name: The name of the language model to use for generating responses.
71
- :param max_tokens: Maximum number of tokens to generate in the response. If None, uses the model's default value.
72
- :param temperature: Temperature to use for the model.
73
- :param use_preview_models: If True, allows the use of preview models from Groq.
72
+ :param chat_model_kwargs: Optional kwargs passed through to ``get_chat_model`` for the main model.
74
73
  :param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
75
74
  :param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
76
75
  :param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
77
76
  :param query_rewriting_model: Name of the model to use for query rewriting.
78
- :param query_rewriting_temperature: Temperature to use for query rewriting.
77
+ :param query_rewriting_chat_model_kwargs: Optional kwargs for the query rewriting model.
78
+ Inheritance order is: ``chat_model_kwargs`` -> ``{"temperature": 0.1}`` defaults -> explicit query rewriting kwargs.
79
79
  :param rerank_model: Name of the Pinecone reranker model to use.
80
80
  :param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
81
81
  :param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
82
82
  :param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
83
- :param retry_config: Optional RetryConfig for Groq API calls
84
- :param fallback_models: Optional list of model names to fall back to if the Groq API fails.
85
83
 
86
84
  :return: A compiled LangGraph with the chosen features
87
85
 
@@ -95,17 +93,18 @@ def advanced_rag(
95
93
  # Parallel searcher for vectorstore(s)
96
94
  searcher = search_runnable(vectorstore)
97
95
 
96
+ primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
97
+
98
98
  # Initialise models
99
- model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config, fallback_models)
99
+ model = get_chat_model(model_name, **primary_model_kwargs)
100
100
 
101
101
  if enable_query_rewriting:
102
- query_rewriting_model_instance = get_chat_model(
103
- model_name=query_rewriting_model,
104
- temperature=query_rewriting_temperature,
105
- use_preview_models=use_preview_models,
106
- retry_config=retry_config,
107
- fallback_models=fallback_models,
108
- )
102
+ query_model_kwargs: ChatModelKwargs = {
103
+ **primary_model_kwargs,
104
+ "temperature": 0.1,
105
+ }
106
+ query_model_kwargs.update(sanitize_chat_model_kwargs(query_rewriting_chat_model_kwargs))
107
+ query_rewriting_model_instance = get_chat_model(model_name=query_rewriting_model, **query_model_kwargs)
109
108
 
110
109
  # --------------- Define functions for the graph nodes ---------------
111
110
 
@@ -209,7 +208,12 @@ def advanced_rag(
209
208
  def generate_answer(state: HybridGraphState) -> HybridGraphState:
210
209
  """Generate a response from the LLM"""
211
210
  # Format the prompt using LangChain template
212
- prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
211
+ prompt_input = {
212
+ "context": combine_documents(state["docs"]),
213
+ "question": state["question"],
214
+ "chat_history": state.get("chat_history", ""),
215
+ }
216
+
213
217
  final_prompt = prompt_template.format_messages(**prompt_input)
214
218
 
215
219
  response = model.invoke(final_prompt)
@@ -274,11 +278,6 @@ if __name__ == "__main__":
274
278
  twiki = get_vectorstore("twiki_prod")
275
279
  mkdocs = get_vectorstore("mkdocs_prod_v1")
276
280
 
277
- retry_config = RetryConfig(
278
- max_retries=1,
279
- max_delay=120.0,
280
- )
281
-
282
281
  # Create the hybrid graph
283
282
  graph = advanced_rag(
284
283
  vectorstore=[twiki, mkdocs],
@@ -1,24 +1,19 @@
1
1
  from operator import itemgetter
2
- from typing import Optional
3
2
 
4
3
  from langchain_core.prompts import ChatPromptTemplate
5
4
  from langchain_core.runnables import RunnableSerializable
6
5
 
7
- from chATLAS_Chains.llm.groq import RetryConfig
8
- from chATLAS_Chains.llm.model_selection import get_chat_model
6
+ from chATLAS_Chains.llm.model_selection import ChatModelKwargs, get_chat_model, sanitize_chat_model_kwargs
9
7
  from chATLAS_Chains.search.basic import search_runnable
10
8
  from chATLAS_Chains.utils.doc_utils import combine_documents
11
- from chATLAS_Embed.Base import VectorStore
9
+ from chATLAS_Embed.VectorStores import VectorStore
12
10
 
13
11
 
14
12
  def basic_retrieval_chain(
15
13
  prompt: str,
16
14
  vectorstore: VectorStore | list[VectorStore],
17
15
  model_name: str,
18
- max_tokens: int | None = None,
19
- temperature: float | None = None,
20
- use_preview_models: bool = False,
21
- retry_config: RetryConfig | None = None,
16
+ chat_model_kwargs: ChatModelKwargs | None = None,
22
17
  ) -> RunnableSerializable:
23
18
  """
24
19
  Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
@@ -29,21 +24,16 @@ def basic_retrieval_chain(
29
24
  :type vectorstore: Any
30
25
  :param model_name: The name of the chat model to use for generating responses.
31
26
  :type model_name: str
32
- :param max_tokens: The maximum number of tokens to generate in the response. Defaults to None.
33
- :type max_tokens: int | None
34
- :param temperature: The temperature to use for the model's response generation. Defaults to None.
35
- :type temperature: float | None
36
- :param use_preview_models: Whether to allow the use of preview models from Groq. Defaults to False.
37
- :type use_preview_models: bool
38
- :param retry_config: Configuration for retrying requests to the model in case of failures. Defaults to None, which will use the default retry configuration.
39
- :type retry_config: RetryConfig | None
27
+ :param chat_model_kwargs: Optional kwargs passed through to ``get_chat_model``.
28
+ :type chat_model_kwargs: ChatModelKwargs | None
40
29
 
41
30
  :return: A LangChain RunnableSerializable chain that performs retrieval and response generation.
42
31
  :rtype: RunnableSerializable
43
32
 
44
33
  """
45
34
  prompt_template = ChatPromptTemplate.from_template(prompt)
46
- model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
35
+ sanitized_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
36
+ model = get_chat_model(model_name, **sanitized_kwargs)
47
37
 
48
38
  search = search_runnable(vectorstore)
49
39
 
@@ -72,7 +62,8 @@ if __name__ == "__main__":
72
62
  chain = basic_retrieval_chain(
73
63
  prompt=CHAT_PROMPT_TEMPLATE,
74
64
  vectorstore=[twiki_vectorstore, mkdocs_vectorstore],
75
- model_name=GROQ_PRODUCTION_MODELS[0],
65
+ model_name="openai/gpt-oss-120b",
66
+ # chat_model_kwargs={"service_provider": "groq"},
76
67
  # model_name="meta-llama/llama-4-maverick-17b-128e-instruct",
77
68
  # model_name="gemma2-9b-it",
78
69
  # model_name="qwen-qwq-32b",
@@ -31,8 +31,6 @@ def basic_retrieval_graph(
31
31
  model_name: str,
32
32
  max_tokens: int | None = None,
33
33
  temperature: float | None = None,
34
- use_preview_models: bool = False,
35
- retry_config: RetryConfig | None = None,
36
34
  ) -> CompiledStateGraph:
37
35
  """
38
36
  Baseline RAG retrieval graph using LangGraph. Searches one or several vectorstores,
@@ -47,7 +45,7 @@ def basic_retrieval_graph(
47
45
  A LangGraph graph that can be executed for RAG
48
46
  """
49
47
  # Initialize the model and prompt template
50
- model = get_chat_model(model_name, use_preview_models=use_preview_models)
48
+ model = get_chat_model(model_name)
51
49
  prompt_template = ChatPromptTemplate.from_template(prompt)
52
50
 
53
51
  # Create a list of retrievers from the vectorstore(s)
@@ -164,9 +162,7 @@ if __name__ == "__main__":
164
162
  # model_name="gemma2-9b-it" # smaller context window, to check the doc truncation
165
163
  # model_name = "gpt-4o-mini"
166
164
 
167
- graph = basic_retrieval_graph(
168
- prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vs, model_name=model_name, use_preview_models=True
169
- )
165
+ graph = basic_retrieval_graph(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vs, model_name=model_name)
170
166
 
171
167
  result = graph.invoke(
172
168
  {