chATLAS_Chains 0.1.5__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.5 → chatlas_chains-0.1.7}/PKG-INFO +78 -7
  2. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/README.md +77 -6
  3. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/advanced.py +56 -60
  4. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/basic.py +9 -18
  5. {chatlas_chains-0.1.5 → 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.5 → chatlas_chains-0.1.7}/chATLAS_Chains/documents/rerank.py +2 -1
  8. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/documents/rrf.py +29 -26
  9. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/llm/groq.py +3 -2
  10. chatlas_chains-0.1.7/chATLAS_Chains/llm/model_selection.py +302 -0
  11. chatlas_chains-0.1.7/chATLAS_Chains/llm/runnables.py +156 -0
  12. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/starters.py +136 -0
  13. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/search/basic.py +1 -1
  14. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/vectorstore.py +4 -3
  15. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/PKG-INFO +78 -7
  16. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/SOURCES.txt +7 -1
  17. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/requires.txt +1 -1
  18. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/pyproject.toml +2 -2
  19. chatlas_chains-0.1.7/tests/conftest.py +99 -0
  20. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/test_chains.py +41 -33
  21. chatlas_chains-0.1.7/tests/test_chat_model_kwargs.py +168 -0
  22. chatlas_chains-0.1.7/tests/test_conversational.py +238 -0
  23. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/test_groq.py +0 -1
  24. chatlas_chains-0.1.7/tests/test_llm_runnables.py +48 -0
  25. chatlas_chains-0.1.7/tests/test_model_selection.py +129 -0
  26. chatlas_chains-0.1.7/tests/test_rrf.py +126 -0
  27. chatlas_chains-0.1.5/chATLAS_Chains/llm/model_selection.py +0 -172
  28. chatlas_chains-0.1.5/tests/conftest.py +0 -68
  29. chatlas_chains-0.1.5/tests/test_llm.py +0 -205
  30. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/LICENSE +0 -0
  31. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/__init__.py +0 -0
  32. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/__init__.py +0 -0
  33. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
  34. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
  35. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/llm/__init__.py +0 -0
  36. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/log.py +0 -0
  37. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/__init__.py +0 -0
  38. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  39. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/query/query_rewriting.py +0 -0
  40. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/search/__init__.py +0 -0
  41. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/utils/__init__.py +0 -0
  42. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains/utils/doc_utils.py +0 -0
  43. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  44. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/chATLAS_Chains.egg-info/top_level.txt +0 -0
  45. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/setup.cfg +0 -0
  46. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/__init__.py +0 -0
  47. {chatlas_chains-0.1.5 → chatlas_chains-0.1.7}/tests/test_search.py +0 -0
  48. {chatlas_chains-0.1.5 → 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.5
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,20 @@ 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
+
159
+ #### 0.1.6
160
+
161
+ Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
162
+
163
+ Add `fallback_models` optional argument to `advanced_rag`
164
+
94
165
  #### 0.1.5
95
166
 
96
167
  Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
@@ -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,20 @@ 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
+
133
+ #### 0.1.6
134
+
135
+ Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
136
+
137
+ Add `fallback_models` optional argument to `advanced_rag`
138
+
68
139
  #### 0.1.5
69
140
 
70
141
  Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
@@ -113,4 +184,4 @@ chATLAS_Benchmark is released under Apache v2.0 license.
113
184
 
114
185
  *For questions and support, please [contact](mailto:joseph.caimin.egan@cern.ch)*
115
186
 
116
- </div>
187
+ </div>
@@ -11,6 +11,7 @@ upweighting results that appear in both
11
11
  """
12
12
 
13
13
  import argparse
14
+ import logging
14
15
  import os
15
16
  import sys
16
17
  from typing import TypedDict
@@ -22,13 +23,19 @@ from langgraph.graph.state import CompiledStateGraph
22
23
 
23
24
  from chATLAS_Chains.documents.rerank import rerank_documents
24
25
  from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
25
- from chATLAS_Chains.llm.groq import RetryConfig
26
- 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
+ )
27
32
  from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
28
33
  from chATLAS_Chains.query.query_rewriting import rewrite_query
29
34
  from chATLAS_Chains.search.basic import search_runnable
30
35
  from chATLAS_Chains.utils.doc_utils import combine_documents
31
- from chATLAS_Embed.Base import VectorStore
36
+ from chATLAS_Embed.VectorStores import VectorStore
37
+
38
+ logger = logging.getLogger(__name__)
32
39
 
33
40
 
34
41
  # Define TypedDict for the simplified state
@@ -37,26 +44,24 @@ class HybridGraphState(TypedDict, total=False):
37
44
  search_kwargs: dict
38
45
  docs: list[Document]
39
46
  answer: str
47
+ chat_history: str # Optional
40
48
 
41
49
 
42
50
  def advanced_rag(
43
51
  vectorstore: VectorStore | list[VectorStore],
44
52
  model_name: str,
45
53
  prompt: str | None = None,
46
- max_tokens: int | None = None,
47
- temperature: float = 0.1,
48
- use_preview_models: bool = False,
54
+ chat_model_kwargs: ChatModelKwargs | None = None,
49
55
  enable_query_rewriting: bool = False,
50
56
  enable_rrf: bool = False,
51
57
  enable_reranking: bool = False,
52
58
  # enable_self_evaluation: bool = False,
53
59
  query_rewriting_model: str = GROQ_PRODUCTION_MODELS[0],
54
- query_rewriting_temperature: float = 0.1,
60
+ query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
55
61
  rerank_model: str = "cohere-rerank-3.5",
56
62
  pinecone_api_key: str | None = None,
57
63
  rrf_constant: float = 60.0,
58
64
  rrf_weights: dict[str, float] | None = None,
59
- retry_config: RetryConfig | None = None,
60
65
  ) -> CompiledStateGraph:
61
66
  """
62
67
  Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
@@ -64,19 +69,17 @@ def advanced_rag(
64
69
  :param prompt: The prompt template to use for the language model. If None, uses chATLAS_Chains.prompt.starters.CHAT_PROMPT_TEMPLATE
65
70
  :param vectorstore: Single vectorstore instance or list of vectorstore instances to search
66
71
  :param model_name: The name of the language model to use for generating responses.
67
- :param max_tokens: Maximum number of tokens to generate in the response. If None, uses the model's default value.
68
- :param temperature: Temperature to use for the model.
69
- :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.
70
73
  :param enable_query_rewriting: Whether to enable LLM-powered query rewriting.
71
74
  :param enable_rrf: Whether to enable RRF (Reciprocal Rank Fusion) for combining results from multiple vectorstores.
72
75
  :param enable_reranking: Whether to rerank the retrieved results using the Pinecone API.
73
76
  :param query_rewriting_model: Name of the model to use for query rewriting.
74
- :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.
75
79
  :param rerank_model: Name of the Pinecone reranker model to use.
76
80
  :param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
77
81
  :param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
78
82
  :param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
79
- :param retry_config: Optional RetryConfig for Groq API calls
80
83
 
81
84
  :return: A compiled LangGraph with the chosen features
82
85
 
@@ -90,15 +93,18 @@ def advanced_rag(
90
93
  # Parallel searcher for vectorstore(s)
91
94
  searcher = search_runnable(vectorstore)
92
95
 
96
+ primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
97
+
93
98
  # Initialise models
94
- model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
99
+ model = get_chat_model(model_name, **primary_model_kwargs)
95
100
 
96
101
  if enable_query_rewriting:
97
- query_rewriting_model_instance = get_chat_model(
98
- model_name=query_rewriting_model,
99
- temperature=query_rewriting_temperature,
100
- use_preview_models=use_preview_models,
101
- )
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)
102
108
 
103
109
  # --------------- Define functions for the graph nodes ---------------
104
110
 
@@ -107,7 +113,7 @@ def advanced_rag(
107
113
  question = state["question"]
108
114
 
109
115
  if not enable_query_rewriting:
110
- print("[WARNING] Query rewriting disabled, but query_rewriting node was called.")
116
+ logger.warning("Query rewriting disabled, but query_rewriting node was called.")
111
117
  # return original state
112
118
  return {**state}
113
119
 
@@ -128,7 +134,7 @@ def advanced_rag(
128
134
 
129
135
  search_kwargs = state.get("search_kwargs", {})
130
136
  if not search_kwargs:
131
- print("[WARNING] No search_kwargs provided, using defaults.")
137
+ logger.warning("No search_kwargs provided, using defaults.")
132
138
  results = searcher.invoke(question)
133
139
  else:
134
140
  results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
@@ -136,7 +142,7 @@ def advanced_rag(
136
142
  if "docs" not in results:
137
143
  raise Exception("Search results missing 'docs' field")
138
144
 
139
- print(f"Retrieved {len(results['docs'])} documents")
145
+ logger.debug(f"Retrieved {len(results['docs'])} documents")
140
146
 
141
147
  return {
142
148
  **state,
@@ -150,7 +156,7 @@ def advanced_rag(
150
156
  """
151
157
 
152
158
  if not enable_reranking:
153
- print("[WARNING] Reranking disabled, but rerank node was called.")
159
+ logger.warning("Reranking disabled, but rerank node was called.")
154
160
  # return original state
155
161
  return {**state}
156
162
 
@@ -166,7 +172,7 @@ def advanced_rag(
166
172
  )
167
173
 
168
174
  except Exception as e:
169
- print(f"[WARNING] Reranking failed with exception {e}. Returning original documents.")
175
+ logger.warning(f"Reranking failed with exception {e}. Returning original documents.")
170
176
  reranked_docs = docs
171
177
 
172
178
  return {
@@ -177,13 +183,13 @@ def advanced_rag(
177
183
  def rrf(state: HybridGraphState) -> HybridGraphState:
178
184
  """Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
179
185
  if not enable_rrf:
180
- print("[WARNING] RRF disabled, but rrf node was called.")
186
+ logger.warning("RRF disabled, but rrf node was called.")
181
187
  # return original state
182
188
  return {**state}
183
189
 
184
190
  docs = state.get("docs", [])
185
191
  if not docs:
186
- print("[WARNING] No documents retrieved for RRF.")
192
+ logger.warning("No documents retrieved for RRF.")
187
193
  return {**state}
188
194
 
189
195
  try:
@@ -191,7 +197,7 @@ def advanced_rag(
191
197
  results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
192
198
  )
193
199
  except Exception as e:
194
- print(f"[WARNING] RRF failed with exception {e}. Returning original documents.")
200
+ logger.warning(f"RRF failed with exception {e}. Returning original documents.")
195
201
  rrf_docs = docs
196
202
 
197
203
  return {
@@ -201,9 +207,13 @@ def advanced_rag(
201
207
 
202
208
  def generate_answer(state: HybridGraphState) -> HybridGraphState:
203
209
  """Generate a response from the LLM"""
204
-
205
210
  # Format the prompt using LangChain template
206
- 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
+
207
217
  final_prompt = prompt_template.format_messages(**prompt_input)
208
218
 
209
219
  response = model.invoke(final_prompt)
@@ -266,42 +276,28 @@ if __name__ == "__main__":
266
276
  from chATLAS_Chains.vectorstore import get_vectorstore
267
277
 
268
278
  twiki = get_vectorstore("twiki_prod")
269
-
270
- retry_config = RetryConfig(
271
- max_retries=1,
272
- max_delay=120.0,
273
- )
279
+ mkdocs = get_vectorstore("mkdocs_prod_v1")
274
280
 
275
281
  # Create the hybrid graph
276
282
  graph = advanced_rag(
277
- vectorstore=[twiki],
283
+ vectorstore=[twiki, mkdocs],
278
284
  model_name=GROQ_PRODUCTION_MODELS[0],
279
285
  enable_query_rewriting=True,
280
286
  enable_rrf=True,
281
- enable_reranking=True,
287
+ enable_reranking=False,
288
+ )
289
+
290
+ ans = graph.invoke(
291
+ {
292
+ "question": "What is the crack veto in electron reconstruction?",
293
+ "search_kwargs": {
294
+ "k_text": 3,
295
+ "k": 15,
296
+ "date_filter": "01-01-2010",
297
+ # "type": ["twiki"],
298
+ },
299
+ }
282
300
  )
283
301
 
284
- # Test query
285
- try:
286
- ans = graph.invoke(
287
- {
288
- "question": "How can one check for and remove bad or corrupted events in the analysis?",
289
- "search_kwargs": {
290
- "k_text": 3,
291
- "k": 15,
292
- "date_filter": "01-01-2010",
293
- # "type": ["twiki"],
294
- },
295
- }
296
- )
297
-
298
- print(f"Number of docs is : {len(ans['docs'])}")
299
- print(f"Answer: {ans['answer']}")
300
-
301
- except Exception as e:
302
- print(f"❌ Graph execution failed with error: {e}")
303
- print(f"Error type: {type(e).__name__}")
304
- import traceback
305
-
306
- traceback.print_exc()
307
- sys.exit(1)
302
+ print(f"Number of docs is : {len(ans['docs'])}")
303
+ print(f"Answer: {ans['answer']}")
@@ -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
  {