chATLAS_Chains 0.1.6__tar.gz → 0.2.0__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 (49) hide show
  1. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/PKG-INFO +72 -7
  2. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/README.md +71 -6
  3. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/advanced.py +66 -54
  4. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/basic.py +9 -18
  5. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/basic_graph.py +2 -6
  6. chatlas_chains-0.2.0/chATLAS_Chains/chains/conversational_graph.py +575 -0
  7. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/documents/rerank.py +2 -1
  8. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/llm/groq.py +113 -4
  9. chatlas_chains-0.2.0/chATLAS_Chains/llm/model_selection.py +302 -0
  10. chatlas_chains-0.2.0/chATLAS_Chains/llm/runnables.py +156 -0
  11. chatlas_chains-0.2.0/chATLAS_Chains/prompt/starters.py +314 -0
  12. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/search/basic.py +1 -1
  13. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/vectorstore.py +4 -3
  14. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/PKG-INFO +72 -7
  15. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/SOURCES.txt +6 -1
  16. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/requires.txt +1 -1
  17. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/pyproject.toml +2 -2
  18. chatlas_chains-0.2.0/tests/conftest.py +99 -0
  19. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_chains.py +11 -15
  20. chatlas_chains-0.2.0/tests/test_chat_model_kwargs.py +170 -0
  21. chatlas_chains-0.2.0/tests/test_conversational.py +238 -0
  22. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_groq.py +23 -5
  23. chatlas_chains-0.2.0/tests/test_llm_runnables.py +48 -0
  24. chatlas_chains-0.2.0/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/chATLAS_Chains/prompt/starters.py +0 -149
  27. chatlas_chains-0.1.6/tests/conftest.py +0 -68
  28. chatlas_chains-0.1.6/tests/test_llm.py +0 -227
  29. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/LICENSE +0 -0
  30. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/__init__.py +0 -0
  31. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/__init__.py +0 -0
  32. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
  33. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
  34. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/documents/rrf.py +0 -0
  35. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/llm/__init__.py +0 -0
  36. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/log.py +0 -0
  37. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/prompt/__init__.py +0 -0
  38. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  39. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/query/query_rewriting.py +0 -0
  40. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/search/__init__.py +0 -0
  41. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/utils/__init__.py +0 -0
  42. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains/utils/doc_utils.py +0 -0
  43. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  44. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/chATLAS_Chains.egg-info/top_level.txt +0 -0
  45. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/setup.cfg +0 -0
  46. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/__init__.py +0 -0
  47. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_rrf.py +0 -0
  48. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/tests/test_search.py +0 -0
  49. {chatlas_chains-0.1.6 → chatlas_chains-0.2.0}/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.2.0
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,25 @@ 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,
65
+ skip_generation: bool = False,
64
66
  ) -> CompiledStateGraph:
65
67
  """
66
68
  Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
@@ -68,20 +70,17 @@ def advanced_rag(
68
70
  :param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
69
71
  :param vectorstore: Single vectorstore instance or list of vectorstore instances to search
70
72
  :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.
73
+ :param chat_model_kwargs: Optional kwargs passed through to ``get_chat_model`` for the main model.
74
74
  :param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
75
75
  :param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
76
76
  :param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
77
77
  :param query_rewriting_model: Name of the model to use for query rewriting.
78
- :param query_rewriting_temperature: Temperature to use for query rewriting.
78
+ :param query_rewriting_chat_model_kwargs: Optional kwargs for the query rewriting model.
79
+ Inheritance order is: ``chat_model_kwargs`` -> ``{"temperature": 0.1}`` defaults -> explicit query rewriting kwargs.
79
80
  :param rerank_model: Name of the Pinecone reranker model to use.
80
81
  :param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
81
82
  :param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
82
83
  :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
84
 
86
85
  :return: A compiled LangGraph with the chosen features
87
86
 
@@ -95,17 +94,18 @@ def advanced_rag(
95
94
  # Parallel searcher for vectorstore(s)
96
95
  searcher = search_runnable(vectorstore)
97
96
 
97
+ primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
98
+
98
99
  # Initialise models
99
- model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config, fallback_models)
100
+ model = get_chat_model(model_name, **primary_model_kwargs)
100
101
 
101
102
  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
- )
103
+ query_model_kwargs: ChatModelKwargs = {
104
+ **primary_model_kwargs,
105
+ "temperature": 0.1,
106
+ }
107
+ query_model_kwargs.update(sanitize_chat_model_kwargs(query_rewriting_chat_model_kwargs))
108
+ query_rewriting_model_instance = get_chat_model(model_name=query_rewriting_model, **query_model_kwargs)
109
109
 
110
110
  # --------------- Define functions for the graph nodes ---------------
111
111
 
@@ -209,7 +209,12 @@ def advanced_rag(
209
209
  def generate_answer(state: HybridGraphState) -> HybridGraphState:
210
210
  """Generate a response from the LLM"""
211
211
  # Format the prompt using LangChain template
212
- prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
212
+ prompt_input = {
213
+ "context": combine_documents(state["docs"]),
214
+ "question": state["question"],
215
+ "chat_history": state.get("chat_history", ""),
216
+ }
217
+
213
218
  final_prompt = prompt_template.format_messages(**prompt_input)
214
219
 
215
220
  response = model.invoke(final_prompt)
@@ -223,50 +228,62 @@ def advanced_rag(
223
228
  # --------------- Build the graph ---------------
224
229
  graph = StateGraph(HybridGraphState)
225
230
 
226
- # Add all the nodes, but don't link to them if not using
227
231
  graph.add_node("query_rewrite", query_rewriting)
228
232
  graph.add_node("retrieval", retrieval)
229
233
  graph.add_node("rrf", rrf)
230
234
  graph.add_node("rerank", rerank)
231
- graph.add_node("generate", generate_answer)
232
- # graph.add_node("assess", assess_answer)
233
- # graph.add_node("refine", refine_answer)
235
+ if not skip_generation:
236
+ graph.add_node("generate", generate_answer)
234
237
 
235
238
  if enable_query_rewriting:
236
- # rewrite the query first
237
239
  graph.add_edge("query_rewrite", "retrieval")
238
240
  graph.set_entry_point("query_rewrite")
239
241
  else:
240
- # start with retrieval
241
242
  graph.set_entry_point("retrieval")
242
243
 
243
- if not enable_rrf and not enable_reranking:
244
- # no document processing, go straight to generation
245
- graph.add_edge("retrieval", "generate")
246
-
247
- elif enable_rrf and not enable_reranking:
248
- # retrieve → rrf → generate
244
+ # Determine the last pre-generation node
245
+ if enable_rrf and enable_reranking:
246
+ last_node = "rerank"
249
247
  graph.add_edge("retrieval", "rrf")
250
- graph.add_edge("rrf", "generate")
251
-
252
- elif not enable_rrf and enable_reranking:
253
- # retrieve → rerank → generate
248
+ graph.add_edge("rrf", "rerank")
249
+ elif enable_rrf:
250
+ last_node = "rrf"
251
+ graph.add_edge("retrieval", "rrf")
252
+ elif enable_reranking:
253
+ last_node = "rerank"
254
254
  graph.add_edge("retrieval", "rerank")
255
- graph.add_edge("rerank", "generate")
256
-
257
255
  else:
258
- # retrieve → rrf → rerank → generate
259
- graph.add_edge("retrieval", "rrf")
260
- graph.add_edge("rrf", "rerank")
261
- graph.add_edge("rerank", "generate")
256
+ last_node = "retrieval"
262
257
 
263
- # generate at the end
264
- graph.add_edge("generate", END)
258
+ if skip_generation:
259
+ graph.add_edge(last_node, END)
260
+ else:
261
+ graph.add_edge(last_node, "generate")
262
+ graph.add_edge("generate", END)
265
263
 
266
- # Compile the graph
267
264
  return graph.compile()
268
265
 
269
266
 
267
+ def build_generation_prompt(prompt: str | None = None) -> ChatPromptTemplate:
268
+ """Build the prompt template used by the generation step."""
269
+ if prompt is None:
270
+ prompt = CHAT_PROMPT_TEMPLATE
271
+ return ChatPromptTemplate.from_template(prompt)
272
+
273
+
274
+ def stream_generate_answer(state: HybridGraphState, model, prompt_template: ChatPromptTemplate):
275
+ """Generator that streams LLM tokens for the generation step."""
276
+ prompt_input = {
277
+ "context": combine_documents(state["docs"]),
278
+ "question": state["question"],
279
+ "chat_history": state.get("chat_history", ""),
280
+ }
281
+ final_prompt = prompt_template.format_messages(**prompt_input)
282
+ for chunk in model.stream(final_prompt):
283
+ if chunk.content:
284
+ yield chunk.content
285
+
286
+
270
287
  if __name__ == "__main__":
271
288
  from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
272
289
  from chATLAS_Chains.vectorstore import get_vectorstore
@@ -274,11 +291,6 @@ if __name__ == "__main__":
274
291
  twiki = get_vectorstore("twiki_prod")
275
292
  mkdocs = get_vectorstore("mkdocs_prod_v1")
276
293
 
277
- retry_config = RetryConfig(
278
- max_retries=1,
279
- max_delay=120.0,
280
- )
281
-
282
294
  # Create the hybrid graph
283
295
  graph = advanced_rag(
284
296
  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
  {