chATLAS_Chains 0.1.7__tar.gz → 0.3.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.
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/PKG-INFO +58 -1
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/README.md +57 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/advanced.py +38 -25
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/conversational_graph.py +202 -10
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/groq.py +112 -3
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/prompt/starters.py +29 -0
- chatlas_chains-0.3.0/chATLAS_Chains/router.py +209 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/PKG-INFO +58 -1
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/SOURCES.txt +3 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/pyproject.toml +1 -1
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_chat_model_kwargs.py +4 -2
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_groq.py +23 -4
- chatlas_chains-0.3.0/tests/test_router.py +230 -0
- chatlas_chains-0.3.0/tests/test_router_litellm.py +40 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/LICENSE +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/basic.py +2 -2
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/basic_graph.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/documents/rerank.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/documents/rrf.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/model_selection.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/runnables.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/query/query_rewriting.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/search/basic.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/utils/doc_utils.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/vectorstore.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/requires.txt +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/top_level.txt +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/setup.cfg +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/__init__.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/conftest.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_chains.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_conversational.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_llm_runnables.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_model_selection.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_rrf.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_search.py +0 -0
- {chatlas_chains-0.1.7 → chatlas_chains-0.3.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.
|
|
3
|
+
Version: 0.3.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
|
|
@@ -78,6 +78,63 @@ More details [here](chATLAS_Chains/chains/README.md)
|
|
|
78
78
|
- `chains.basic.basic_retrieval_chain`
|
|
79
79
|
- `chains.advanced.advanced_rag`
|
|
80
80
|
|
|
81
|
+
## Prompt Routing
|
|
82
|
+
|
|
83
|
+
`chATLAS_Chains.router.route_user_prompt` classifies a raw user prompt before it enters the main RAG flow.
|
|
84
|
+
The router defaults to CERN LiteLLM and returns a validated `RouterDecision` with:
|
|
85
|
+
|
|
86
|
+
- `route`: one of `hep_atlas`, `person_lookup`, `dangerous`, `out_of_scope`,
|
|
87
|
+
or `unknown`
|
|
88
|
+
- `confidence`: score from 0 to 1
|
|
89
|
+
- `normalized_query`: cleaned prompt for downstream use
|
|
90
|
+
- `search_kwargs`: optional search parameters
|
|
91
|
+
- `metadata`: route-specific parameters
|
|
92
|
+
|
|
93
|
+
For implementation details and a function-by-function reference, see the
|
|
94
|
+
frontend reviewer document:
|
|
95
|
+
[`../chATLAS_Frontend/docs/INTENT_ROUTER.md`](../chATLAS_Frontend/docs/INTENT_ROUTER.md).
|
|
96
|
+
|
|
97
|
+
```python
|
|
98
|
+
from chATLAS_Chains.router import route_user_prompt
|
|
99
|
+
|
|
100
|
+
decision = route_user_prompt(
|
|
101
|
+
"Who is Jane Doe?",
|
|
102
|
+
chat_model_kwargs={
|
|
103
|
+
"service_provider": "litellm",
|
|
104
|
+
"proxy": "socks5h://localhost:1080", # optional when outside CERN
|
|
105
|
+
},
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
if decision.route == "dangerous":
|
|
109
|
+
print("Apply the safety response policy")
|
|
110
|
+
elif decision.route == "person_lookup":
|
|
111
|
+
print("Use the person lookup flow")
|
|
112
|
+
elif decision.route == "hep_atlas":
|
|
113
|
+
print("Use the HEP/ATLAS search flow")
|
|
114
|
+
elif decision.route == "out_of_scope":
|
|
115
|
+
print("Explain the chATLAS scope")
|
|
116
|
+
else:
|
|
117
|
+
print("Ask for clarification")
|
|
118
|
+
```
|
|
119
|
+
|
|
120
|
+
The dedicated `mc_request` route is currently disabled until the MC workflow is
|
|
121
|
+
ready. MC production and job-option prompts are classified as `hep_atlas` and
|
|
122
|
+
continue through the normal selected workflow.
|
|
123
|
+
|
|
124
|
+
Router context uses a hybrid policy: the current user prompt is authoritative,
|
|
125
|
+
and recent conversation context is only a bounded disambiguation aid for
|
|
126
|
+
follow-ups such as "What about Run 3?". The router keeps recent user prompts,
|
|
127
|
+
includes only short assistant messages, and drops long assistant answers rather
|
|
128
|
+
than summarizing them.
|
|
129
|
+
|
|
130
|
+
For live LiteLLM calls, set:
|
|
131
|
+
|
|
132
|
+
```sh
|
|
133
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your litellm key"
|
|
134
|
+
```
|
|
135
|
+
|
|
136
|
+
If LiteLLM is unavailable or returns invalid JSON, the router falls back to deterministic rules by default.
|
|
137
|
+
|
|
81
138
|
### Model Configuration in Chains
|
|
82
139
|
|
|
83
140
|
Supported chain constructors now accept a typed `chat_model_kwargs` argument for model options (for example:
|
|
@@ -52,6 +52,63 @@ More details [here](chATLAS_Chains/chains/README.md)
|
|
|
52
52
|
- `chains.basic.basic_retrieval_chain`
|
|
53
53
|
- `chains.advanced.advanced_rag`
|
|
54
54
|
|
|
55
|
+
## Prompt Routing
|
|
56
|
+
|
|
57
|
+
`chATLAS_Chains.router.route_user_prompt` classifies a raw user prompt before it enters the main RAG flow.
|
|
58
|
+
The router defaults to CERN LiteLLM and returns a validated `RouterDecision` with:
|
|
59
|
+
|
|
60
|
+
- `route`: one of `hep_atlas`, `person_lookup`, `dangerous`, `out_of_scope`,
|
|
61
|
+
or `unknown`
|
|
62
|
+
- `confidence`: score from 0 to 1
|
|
63
|
+
- `normalized_query`: cleaned prompt for downstream use
|
|
64
|
+
- `search_kwargs`: optional search parameters
|
|
65
|
+
- `metadata`: route-specific parameters
|
|
66
|
+
|
|
67
|
+
For implementation details and a function-by-function reference, see the
|
|
68
|
+
frontend reviewer document:
|
|
69
|
+
[`../chATLAS_Frontend/docs/INTENT_ROUTER.md`](../chATLAS_Frontend/docs/INTENT_ROUTER.md).
|
|
70
|
+
|
|
71
|
+
```python
|
|
72
|
+
from chATLAS_Chains.router import route_user_prompt
|
|
73
|
+
|
|
74
|
+
decision = route_user_prompt(
|
|
75
|
+
"Who is Jane Doe?",
|
|
76
|
+
chat_model_kwargs={
|
|
77
|
+
"service_provider": "litellm",
|
|
78
|
+
"proxy": "socks5h://localhost:1080", # optional when outside CERN
|
|
79
|
+
},
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
if decision.route == "dangerous":
|
|
83
|
+
print("Apply the safety response policy")
|
|
84
|
+
elif decision.route == "person_lookup":
|
|
85
|
+
print("Use the person lookup flow")
|
|
86
|
+
elif decision.route == "hep_atlas":
|
|
87
|
+
print("Use the HEP/ATLAS search flow")
|
|
88
|
+
elif decision.route == "out_of_scope":
|
|
89
|
+
print("Explain the chATLAS scope")
|
|
90
|
+
else:
|
|
91
|
+
print("Ask for clarification")
|
|
92
|
+
```
|
|
93
|
+
|
|
94
|
+
The dedicated `mc_request` route is currently disabled until the MC workflow is
|
|
95
|
+
ready. MC production and job-option prompts are classified as `hep_atlas` and
|
|
96
|
+
continue through the normal selected workflow.
|
|
97
|
+
|
|
98
|
+
Router context uses a hybrid policy: the current user prompt is authoritative,
|
|
99
|
+
and recent conversation context is only a bounded disambiguation aid for
|
|
100
|
+
follow-ups such as "What about Run 3?". The router keeps recent user prompts,
|
|
101
|
+
includes only short assistant messages, and drops long assistant answers rather
|
|
102
|
+
than summarizing them.
|
|
103
|
+
|
|
104
|
+
For live LiteLLM calls, set:
|
|
105
|
+
|
|
106
|
+
```sh
|
|
107
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your litellm key"
|
|
108
|
+
```
|
|
109
|
+
|
|
110
|
+
If LiteLLM is unavailable or returns invalid JSON, the router falls back to deterministic rules by default.
|
|
111
|
+
|
|
55
112
|
### Model Configuration in Chains
|
|
56
113
|
|
|
57
114
|
Supported chain constructors now accept a typed `chat_model_kwargs` argument for model options (for example:
|
|
@@ -62,6 +62,7 @@ def advanced_rag(
|
|
|
62
62
|
pinecone_api_key: str | None = None,
|
|
63
63
|
rrf_constant: float = 60.0,
|
|
64
64
|
rrf_weights: dict[str, float] | None = None,
|
|
65
|
+
skip_generation: bool = False,
|
|
65
66
|
) -> CompiledStateGraph:
|
|
66
67
|
"""
|
|
67
68
|
Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
|
|
@@ -227,50 +228,62 @@ def advanced_rag(
|
|
|
227
228
|
# --------------- Build the graph ---------------
|
|
228
229
|
graph = StateGraph(HybridGraphState)
|
|
229
230
|
|
|
230
|
-
# Add all the nodes, but don't link to them if not using
|
|
231
231
|
graph.add_node("query_rewrite", query_rewriting)
|
|
232
232
|
graph.add_node("retrieval", retrieval)
|
|
233
233
|
graph.add_node("rrf", rrf)
|
|
234
234
|
graph.add_node("rerank", rerank)
|
|
235
|
-
|
|
236
|
-
|
|
237
|
-
# graph.add_node("refine", refine_answer)
|
|
235
|
+
if not skip_generation:
|
|
236
|
+
graph.add_node("generate", generate_answer)
|
|
238
237
|
|
|
239
238
|
if enable_query_rewriting:
|
|
240
|
-
# rewrite the query first
|
|
241
239
|
graph.add_edge("query_rewrite", "retrieval")
|
|
242
240
|
graph.set_entry_point("query_rewrite")
|
|
243
241
|
else:
|
|
244
|
-
# start with retrieval
|
|
245
242
|
graph.set_entry_point("retrieval")
|
|
246
243
|
|
|
247
|
-
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
elif enable_rrf and not enable_reranking:
|
|
252
|
-
# retrieve → rrf → generate
|
|
244
|
+
# Determine the last pre-generation node
|
|
245
|
+
if enable_rrf and enable_reranking:
|
|
246
|
+
last_node = "rerank"
|
|
253
247
|
graph.add_edge("retrieval", "rrf")
|
|
254
|
-
graph.add_edge("rrf", "
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
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"
|
|
258
254
|
graph.add_edge("retrieval", "rerank")
|
|
259
|
-
graph.add_edge("rerank", "generate")
|
|
260
|
-
|
|
261
255
|
else:
|
|
262
|
-
|
|
263
|
-
graph.add_edge("retrieval", "rrf")
|
|
264
|
-
graph.add_edge("rrf", "rerank")
|
|
265
|
-
graph.add_edge("rerank", "generate")
|
|
256
|
+
last_node = "retrieval"
|
|
266
257
|
|
|
267
|
-
|
|
268
|
-
|
|
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)
|
|
269
263
|
|
|
270
|
-
# Compile the graph
|
|
271
264
|
return graph.compile()
|
|
272
265
|
|
|
273
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
|
+
|
|
274
287
|
if __name__ == "__main__":
|
|
275
288
|
from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
|
|
276
289
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
@@ -16,7 +16,7 @@ from langgraph.graph.message import add_messages
|
|
|
16
16
|
from langgraph.graph.state import CompiledStateGraph
|
|
17
17
|
from sentence_transformers.util import cos_sim
|
|
18
18
|
|
|
19
|
-
from chATLAS_Chains.chains.advanced import advanced_rag
|
|
19
|
+
from chATLAS_Chains.chains.advanced import advanced_rag, build_generation_prompt, stream_generate_answer
|
|
20
20
|
from chATLAS_Chains.llm.model_selection import (
|
|
21
21
|
GROQ_PRODUCTION_MODELS,
|
|
22
22
|
ChatModelKwargs,
|
|
@@ -27,6 +27,7 @@ from chATLAS_Chains.prompt.starters import (
|
|
|
27
27
|
CHAT_PROMPT_TEMPLATE,
|
|
28
28
|
CONVERSATIONAL_CHAT_PROMPT_TEMPLATE,
|
|
29
29
|
QUERY_CONTEXTUALIZATION_PROMPT,
|
|
30
|
+
STREAMING_PRODUCTION_PROMPT_TEMPLATE,
|
|
30
31
|
)
|
|
31
32
|
|
|
32
33
|
logger = logging.getLogger(__name__)
|
|
@@ -176,21 +177,19 @@ def conversational_retrieval_graph(
|
|
|
176
177
|
prompt = CONVERSATIONAL_CHAT_PROMPT_TEMPLATE
|
|
177
178
|
|
|
178
179
|
if contextualization_model is None:
|
|
179
|
-
contextualization_model =
|
|
180
|
+
contextualization_model = "gpt-5-nano"
|
|
181
|
+
# gpt-5 family does not support max_tokens or temperature params;
|
|
182
|
+
# override any inherited kwargs to avoid silent empty responses
|
|
183
|
+
contextualization_chat_model_kwargs = {"service_provider": "openai"}
|
|
180
184
|
|
|
181
185
|
primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
|
|
182
186
|
|
|
183
187
|
contextualization_llm = None
|
|
184
188
|
if enable_query_contextualization:
|
|
185
|
-
|
|
186
|
-
**primary_model_kwargs,
|
|
187
|
-
"max_tokens": 256,
|
|
188
|
-
"temperature": 0.1,
|
|
189
|
-
}
|
|
190
|
-
contextualization_kwargs.update(sanitize_chat_model_kwargs(contextualization_chat_model_kwargs))
|
|
189
|
+
ctx_kwargs = sanitize_chat_model_kwargs(contextualization_chat_model_kwargs)
|
|
191
190
|
contextualization_llm = get_chat_model(
|
|
192
191
|
model_name=contextualization_model,
|
|
193
|
-
**
|
|
192
|
+
**ctx_kwargs,
|
|
194
193
|
)
|
|
195
194
|
contextualization_prompt = ChatPromptTemplate.from_template(QUERY_CONTEXTUALIZATION_PROMPT)
|
|
196
195
|
|
|
@@ -235,7 +234,9 @@ def conversational_retrieval_graph(
|
|
|
235
234
|
try:
|
|
236
235
|
formatted_prompt = contextualization_prompt.format(chat_history=chat_history, question=question)
|
|
237
236
|
response = contextualization_llm.invoke(formatted_prompt)
|
|
238
|
-
|
|
237
|
+
rewritten = str(response.content).strip()
|
|
238
|
+
# Model sometimes returns empty when it means "unchanged" — fall back to original
|
|
239
|
+
standalone_question = rewritten if rewritten else question
|
|
239
240
|
|
|
240
241
|
if standalone_question != question:
|
|
241
242
|
logger.info(f"Contextualized question: '{question}' -> '{standalone_question}'")
|
|
@@ -323,6 +324,197 @@ def conversational_retrieval_graph(
|
|
|
323
324
|
return graph.compile(checkpointer=checkpointer)
|
|
324
325
|
|
|
325
326
|
|
|
327
|
+
def conversational_retrieval_streaming(
|
|
328
|
+
vectorstore,
|
|
329
|
+
model_name: str,
|
|
330
|
+
checkpointer: BaseCheckpointSaver,
|
|
331
|
+
thread_id: str,
|
|
332
|
+
user_input: str,
|
|
333
|
+
search_kwargs: dict,
|
|
334
|
+
prompt: str | None = None,
|
|
335
|
+
chat_model_kwargs: ChatModelKwargs | None = None,
|
|
336
|
+
enable_query_rewriting: bool = False,
|
|
337
|
+
enable_rrf: bool = True,
|
|
338
|
+
enable_reranking: bool = False,
|
|
339
|
+
enable_query_contextualization: bool = True,
|
|
340
|
+
query_rewriting_model: str | None = None,
|
|
341
|
+
query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
342
|
+
contextualization_model: str | None = None,
|
|
343
|
+
contextualization_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
344
|
+
max_turns: int = 10,
|
|
345
|
+
rerank_model: str = "cohere-rerank-3.5",
|
|
346
|
+
pinecone_api_key: str | None = None,
|
|
347
|
+
rrf_constant: float = 60.0,
|
|
348
|
+
rrf_weights: dict[str, float] | None = None,
|
|
349
|
+
embedder=None,
|
|
350
|
+
source_names: list[str] | None = None,
|
|
351
|
+
):
|
|
352
|
+
"""Streaming variant of the conversational retrieval graph.
|
|
353
|
+
|
|
354
|
+
Generator that yields status dicts during pipeline execution,
|
|
355
|
+
then yields the final ``(docs, StreamingSession)`` tuple.
|
|
356
|
+
"""
|
|
357
|
+
if prompt is None:
|
|
358
|
+
prompt = STREAMING_PRODUCTION_PROMPT_TEMPLATE
|
|
359
|
+
if query_rewriting_model is None:
|
|
360
|
+
query_rewriting_model = model_name
|
|
361
|
+
if contextualization_model is None:
|
|
362
|
+
contextualization_model = "gpt-5-nano"
|
|
363
|
+
# gpt-5 family does not support max_tokens or temperature params;
|
|
364
|
+
# override any inherited kwargs to avoid silent empty responses
|
|
365
|
+
contextualization_chat_model_kwargs = {"service_provider": "openai"}
|
|
366
|
+
|
|
367
|
+
primary_model_kwargs = sanitize_chat_model_kwargs(chat_model_kwargs)
|
|
368
|
+
|
|
369
|
+
conv_graph = conversational_retrieval_graph(
|
|
370
|
+
vectorstore=vectorstore,
|
|
371
|
+
model_name=model_name,
|
|
372
|
+
checkpointer=checkpointer,
|
|
373
|
+
prompt=prompt,
|
|
374
|
+
chat_model_kwargs=chat_model_kwargs,
|
|
375
|
+
enable_query_rewriting=enable_query_rewriting,
|
|
376
|
+
enable_rrf=enable_rrf,
|
|
377
|
+
enable_reranking=enable_reranking,
|
|
378
|
+
enable_query_contextualization=enable_query_contextualization,
|
|
379
|
+
query_rewriting_model=query_rewriting_model,
|
|
380
|
+
query_rewriting_chat_model_kwargs=query_rewriting_chat_model_kwargs,
|
|
381
|
+
contextualization_model=contextualization_model,
|
|
382
|
+
contextualization_chat_model_kwargs=contextualization_chat_model_kwargs,
|
|
383
|
+
max_turns=max_turns,
|
|
384
|
+
rerank_model=rerank_model,
|
|
385
|
+
pinecone_api_key=pinecone_api_key,
|
|
386
|
+
rrf_constant=rrf_constant,
|
|
387
|
+
rrf_weights=rrf_weights,
|
|
388
|
+
embedder=embedder,
|
|
389
|
+
)
|
|
390
|
+
|
|
391
|
+
config = {"configurable": {"thread_id": thread_id}}
|
|
392
|
+
existing_state = conv_graph.get_state(config)
|
|
393
|
+
existing_messages = existing_state.values.get("messages", []) if existing_state.values else []
|
|
394
|
+
|
|
395
|
+
if max_turns is not None and len(existing_messages) >= max_turns * 2:
|
|
396
|
+
limit_msg = f"Conversation turn limit ({max_turns}) reached. Please start a new thread."
|
|
397
|
+
conv_graph.update_state(config, {"messages": [HumanMessage(content=user_input), AIMessage(content=limit_msg)]})
|
|
398
|
+
yield ([], _LimitReachedSession(limit_msg))
|
|
399
|
+
return
|
|
400
|
+
|
|
401
|
+
all_messages = [*existing_messages, HumanMessage(content=user_input)]
|
|
402
|
+
|
|
403
|
+
standalone_question = user_input
|
|
404
|
+
if enable_query_contextualization:
|
|
405
|
+
yield {"type": "status", "stage": "contextualizing", "message": "Analyzing query..."}
|
|
406
|
+
|
|
407
|
+
ctx_kwargs = sanitize_chat_model_kwargs(contextualization_chat_model_kwargs)
|
|
408
|
+
ctx_llm = get_chat_model(model_name=contextualization_model, **ctx_kwargs)
|
|
409
|
+
ctx_prompt = ChatPromptTemplate.from_template(QUERY_CONTEXTUALIZATION_PROMPT)
|
|
410
|
+
|
|
411
|
+
history_messages = [m for m in all_messages if m.content != user_input]
|
|
412
|
+
if history_messages and embedder is not None:
|
|
413
|
+
history_messages = select_relevant_history(
|
|
414
|
+
user_input, history_messages, embedder, max_messages=max_turns * 2
|
|
415
|
+
)
|
|
416
|
+
chat_history = (
|
|
417
|
+
format_chat_history(history_messages, max_turns=max_turns)
|
|
418
|
+
if history_messages
|
|
419
|
+
else "(No previous conversation)"
|
|
420
|
+
)
|
|
421
|
+
|
|
422
|
+
try:
|
|
423
|
+
formatted = ctx_prompt.format(chat_history=chat_history, question=user_input)
|
|
424
|
+
response = ctx_llm.invoke(formatted)
|
|
425
|
+
rewritten = str(response.content).strip()
|
|
426
|
+
# Model sometimes returns empty when it means "unchanged" — fall back to original
|
|
427
|
+
standalone_question = rewritten if rewritten else user_input
|
|
428
|
+
if standalone_question != user_input:
|
|
429
|
+
logger.info(f"Contextualized question: '{user_input}' -> '{standalone_question}'")
|
|
430
|
+
yield {"type": "status", "stage": "contextualizing", "message": f"Searching for: {standalone_question}"}
|
|
431
|
+
except Exception as e:
|
|
432
|
+
logger.warning(f"Query contextualization failed: {e}. Using original question.")
|
|
433
|
+
|
|
434
|
+
sources_label = ", ".join(source_names) if source_names else "documents"
|
|
435
|
+
yield {"type": "status", "stage": "searching", "message": f"Searching {sources_label}..."}
|
|
436
|
+
|
|
437
|
+
retrieval_graph = advanced_rag(
|
|
438
|
+
vectorstore=vectorstore,
|
|
439
|
+
model_name=model_name,
|
|
440
|
+
prompt=prompt,
|
|
441
|
+
chat_model_kwargs=chat_model_kwargs,
|
|
442
|
+
enable_query_rewriting=enable_query_rewriting,
|
|
443
|
+
enable_rrf=enable_rrf,
|
|
444
|
+
enable_reranking=enable_reranking,
|
|
445
|
+
query_rewriting_model=query_rewriting_model,
|
|
446
|
+
query_rewriting_chat_model_kwargs=query_rewriting_chat_model_kwargs,
|
|
447
|
+
rerank_model=rerank_model,
|
|
448
|
+
pinecone_api_key=pinecone_api_key,
|
|
449
|
+
rrf_constant=rrf_constant,
|
|
450
|
+
rrf_weights=rrf_weights,
|
|
451
|
+
skip_generation=True,
|
|
452
|
+
)
|
|
453
|
+
|
|
454
|
+
history_for_generation = format_chat_history(
|
|
455
|
+
[m for m in all_messages if m.content != user_input], max_turns=max_turns
|
|
456
|
+
)
|
|
457
|
+
retrieval_result = retrieval_graph.invoke(
|
|
458
|
+
{
|
|
459
|
+
"question": standalone_question,
|
|
460
|
+
"search_kwargs": search_kwargs,
|
|
461
|
+
"chat_history": history_for_generation,
|
|
462
|
+
}
|
|
463
|
+
)
|
|
464
|
+
docs = retrieval_result.get("docs", [])
|
|
465
|
+
|
|
466
|
+
yield {"type": "status", "stage": "generating", "message": "Thinking..."}
|
|
467
|
+
|
|
468
|
+
streaming_model = get_chat_model(model_name, **primary_model_kwargs)
|
|
469
|
+
prompt_template = build_generation_prompt(prompt)
|
|
470
|
+
|
|
471
|
+
session = _StreamingSession(
|
|
472
|
+
retrieval_state=retrieval_result,
|
|
473
|
+
model=streaming_model,
|
|
474
|
+
prompt_template=prompt_template,
|
|
475
|
+
conv_graph=conv_graph,
|
|
476
|
+
config=config,
|
|
477
|
+
user_input=user_input,
|
|
478
|
+
)
|
|
479
|
+
yield (docs, session)
|
|
480
|
+
|
|
481
|
+
|
|
482
|
+
class _StreamingSession:
|
|
483
|
+
def __init__(self, retrieval_state, model, prompt_template, conv_graph, config, user_input):
|
|
484
|
+
self._retrieval_state = retrieval_state
|
|
485
|
+
self._model = model
|
|
486
|
+
self._prompt_template = prompt_template
|
|
487
|
+
self._conv_graph = conv_graph
|
|
488
|
+
self._config = config
|
|
489
|
+
self._user_input = user_input
|
|
490
|
+
self._full_text = ""
|
|
491
|
+
|
|
492
|
+
def stream_tokens(self):
|
|
493
|
+
for token in stream_generate_answer(self._retrieval_state, self._model, self._prompt_template):
|
|
494
|
+
self._full_text += token
|
|
495
|
+
yield token
|
|
496
|
+
|
|
497
|
+
def finalize(self):
|
|
498
|
+
self._conv_graph.update_state(
|
|
499
|
+
self._config,
|
|
500
|
+
{
|
|
501
|
+
"messages": [HumanMessage(content=self._user_input), AIMessage(content=self._full_text)],
|
|
502
|
+
},
|
|
503
|
+
)
|
|
504
|
+
return self._full_text
|
|
505
|
+
|
|
506
|
+
|
|
507
|
+
class _LimitReachedSession:
|
|
508
|
+
def __init__(self, message):
|
|
509
|
+
self._message = message
|
|
510
|
+
|
|
511
|
+
def stream_tokens(self):
|
|
512
|
+
yield self._message
|
|
513
|
+
|
|
514
|
+
def finalize(self):
|
|
515
|
+
return self._message
|
|
516
|
+
|
|
517
|
+
|
|
326
518
|
if __name__ == "__main__":
|
|
327
519
|
from langchain_core.messages import HumanMessage
|
|
328
520
|
from langgraph.checkpoint.memory import InMemorySaver
|
|
@@ -11,7 +11,8 @@ from typing import Any, Optional
|
|
|
11
11
|
import requests
|
|
12
12
|
from langchain_core.language_models.chat_models import BaseChatModel
|
|
13
13
|
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage
|
|
14
|
-
from langchain_core.
|
|
14
|
+
from langchain_core.messages.ai import AIMessageChunk
|
|
15
|
+
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult
|
|
15
16
|
from pydantic import BaseModel, Field, SecretStr
|
|
16
17
|
from requests.exceptions import ConnectionError, RequestException, Timeout
|
|
17
18
|
|
|
@@ -545,8 +546,116 @@ class AccGPTChatGroq(BaseChatModel, BaseModel):
|
|
|
545
546
|
logger.info("Retryable error occurred, trying next model...")
|
|
546
547
|
continue
|
|
547
548
|
|
|
548
|
-
def _stream(self, **kwargs):
|
|
549
|
-
|
|
549
|
+
def _stream(self, messages: list[BaseMessage], **kwargs):
|
|
550
|
+
"""Stream tokens from the Groq API with retry and fallback support."""
|
|
551
|
+
models_to_try = [self.model_name, *self.fallback_models]
|
|
552
|
+
|
|
553
|
+
for i, model in enumerate(models_to_try):
|
|
554
|
+
try:
|
|
555
|
+
logger.info(f"Streaming with model {i + 1}/{len(models_to_try)}: {model}")
|
|
556
|
+
yield from self._make_streaming_request(messages, model)
|
|
557
|
+
return
|
|
558
|
+
except GroqAPIError as e:
|
|
559
|
+
logger.error(f"Streaming model {model} failed: {e!s}")
|
|
560
|
+
if i == len(models_to_try) - 1:
|
|
561
|
+
raise e
|
|
562
|
+
logger.info("Trying next model...")
|
|
563
|
+
|
|
564
|
+
def _make_streaming_request(self, messages: list[BaseMessage], model_name: str):
|
|
565
|
+
"""Make a single streaming API request with retry logic."""
|
|
566
|
+
formatted_messages = []
|
|
567
|
+
for message in messages:
|
|
568
|
+
if isinstance(message, SystemMessage):
|
|
569
|
+
formatted_messages.append({"role": "system", "content": message.content})
|
|
570
|
+
elif isinstance(message, HumanMessage):
|
|
571
|
+
formatted_messages.append({"role": "user", "content": message.content})
|
|
572
|
+
elif isinstance(message, AIMessage):
|
|
573
|
+
formatted_messages.append({"role": "assistant", "content": message.content})
|
|
574
|
+
|
|
575
|
+
chat_payload = {"messages": formatted_messages, "model": model_name, "n": 1, "stream": True}
|
|
576
|
+
if self.max_tokens is not None:
|
|
577
|
+
chat_payload["max_tokens"] = self.max_tokens
|
|
578
|
+
if self.temperature is not None:
|
|
579
|
+
chat_payload["temperature"] = self.temperature
|
|
580
|
+
|
|
581
|
+
url = f"{self.base_url.get_secret_value()}/chat"
|
|
582
|
+
headers = {"X-API-Key": self.api_key.get_secret_value()}
|
|
583
|
+
|
|
584
|
+
last_exception = None
|
|
585
|
+
for attempt in range(self.retry_config.max_retries + 1):
|
|
586
|
+
try:
|
|
587
|
+
if attempt > 0:
|
|
588
|
+
rate_limit_info = getattr(last_exception, "rate_limit_info", None) if last_exception else None
|
|
589
|
+
delay = self._calculate_delay(attempt - 1, rate_limit_info)
|
|
590
|
+
if delay > 0:
|
|
591
|
+
logger.info(
|
|
592
|
+
f"Stream retry {attempt + 1}/{self.retry_config.max_retries + 1} after {delay:.2f}s"
|
|
593
|
+
)
|
|
594
|
+
time.sleep(delay)
|
|
595
|
+
|
|
596
|
+
response = requests.post(url, json=chat_payload, headers=headers, timeout=self.timeout, stream=True)
|
|
597
|
+
|
|
598
|
+
if response.status_code != 200:
|
|
599
|
+
try:
|
|
600
|
+
result = response.json()
|
|
601
|
+
except json.JSONDecodeError:
|
|
602
|
+
result = {"error": response.text[:500]}
|
|
603
|
+
error_details, is_retryable, rate_limit_info = self._analyze_response_error(response, result)
|
|
604
|
+
error_msg = f"Streaming request failed with status {response.status_code}: {error_details}"
|
|
605
|
+
last_exception = GroqAPIError(
|
|
606
|
+
error_msg,
|
|
607
|
+
response_data=result,
|
|
608
|
+
status_code=response.status_code,
|
|
609
|
+
is_retryable=is_retryable,
|
|
610
|
+
rate_limit_info=rate_limit_info,
|
|
611
|
+
)
|
|
612
|
+
if not is_retryable or attempt == self.retry_config.max_retries:
|
|
613
|
+
raise last_exception
|
|
614
|
+
continue
|
|
615
|
+
|
|
616
|
+
# Stream tokens from SSE response
|
|
617
|
+
for line in response.iter_lines(decode_unicode=True):
|
|
618
|
+
if not line:
|
|
619
|
+
continue
|
|
620
|
+
if not line.startswith("data: "):
|
|
621
|
+
continue
|
|
622
|
+
data_str = line[6:]
|
|
623
|
+
if data_str.strip() == "[DONE]":
|
|
624
|
+
return
|
|
625
|
+
try:
|
|
626
|
+
chunk_data = json.loads(data_str)
|
|
627
|
+
except json.JSONDecodeError:
|
|
628
|
+
continue
|
|
629
|
+
|
|
630
|
+
choices = chunk_data.get("choices", [])
|
|
631
|
+
if not choices:
|
|
632
|
+
continue
|
|
633
|
+
delta = choices[0].get("delta", {})
|
|
634
|
+
content = delta.get("content", "")
|
|
635
|
+
if content:
|
|
636
|
+
yield ChatGenerationChunk(message=AIMessageChunk(content=content))
|
|
637
|
+
|
|
638
|
+
if choices[0].get("finish_reason") == "length":
|
|
639
|
+
logger.warning("Streaming: model ran out of tokens before completing response")
|
|
640
|
+
return
|
|
641
|
+
|
|
642
|
+
except (Timeout, ConnectionError, RequestException) as e:
|
|
643
|
+
error_msg = f"Stream network error: {e!s}"
|
|
644
|
+
logger.error(error_msg)
|
|
645
|
+
last_exception = GroqAPIError(error_msg, is_retryable=True)
|
|
646
|
+
if attempt == self.retry_config.max_retries:
|
|
647
|
+
raise last_exception
|
|
648
|
+
continue
|
|
649
|
+
|
|
650
|
+
except GroqAPIError as e:
|
|
651
|
+
last_exception = e
|
|
652
|
+
if not e.is_retryable or attempt == self.retry_config.max_retries:
|
|
653
|
+
raise e
|
|
654
|
+
continue
|
|
655
|
+
|
|
656
|
+
if last_exception:
|
|
657
|
+
raise last_exception
|
|
658
|
+
raise GroqAPIError("All streaming retry attempts failed")
|
|
550
659
|
|
|
551
660
|
@property
|
|
552
661
|
def _llm_type(self) -> str:
|
|
@@ -185,6 +185,35 @@ Provide your response prioritizing ATLAS internal sources:
|
|
|
185
185
|
}}
|
|
186
186
|
"""
|
|
187
187
|
|
|
188
|
+
STREAMING_PRODUCTION_PROMPT_TEMPLATE = """
|
|
189
|
+
You are an elite-level physicist specializing in high-energy physics, specifically within the ATLAS collaboration at CERN.
|
|
190
|
+
Your role is to serve as a highly reliable expert assistant, providing precise, accurate, and contextually relevant answers to professional physicists' inquiries.
|
|
191
|
+
|
|
192
|
+
### Instructions:
|
|
193
|
+
- **Use Only Provided Context:** You must strictly base your answers on the supplied context. Do not speculate, hallucinate, or fabricate information.
|
|
194
|
+
- **Answer with Clarity and Precision:** Provide clear, concise, and technically accurate responses. Use professional and scientific language appropriate for CERN physicists.
|
|
195
|
+
- **Format and Cite Properly:**
|
|
196
|
+
- Use Markdown for structured formatting (e.g., `**bold**` for emphasis, lists for steps, and code blocks for technical content).
|
|
197
|
+
- When citing a source, reference it inline using the exact document name in square brackets, e.g. [DocumentName].
|
|
198
|
+
- **Show Reasoning and Steps:** When applicable, break down complex answers into clear, logical steps.
|
|
199
|
+
- **Handle Ambiguity:** If the question or context is ambiguous, specify the missing details and suggest clarifying questions.
|
|
200
|
+
- **Prioritize Caution and Verifiability:** When the context is limited, explicitly state the lack of sufficient information rather than making unsupported inferences.
|
|
201
|
+
- **Context Prioritization:** When multiple sources provide relevant information, prioritize: (1) official ATLAS documentation, (2) peer-reviewed publications, (3) technical notes, (4) meeting minutes.
|
|
202
|
+
- **Temporal Awareness:** Be attentive to publication dates in the provided context.
|
|
203
|
+
- **Technical Terminology:** Maintain consistency with ATLAS-specific conventions.
|
|
204
|
+
|
|
205
|
+
### Conversation History:
|
|
206
|
+
{chat_history}
|
|
207
|
+
|
|
208
|
+
### Context (from ATLAS documentation and wiki pages):
|
|
209
|
+
{context}
|
|
210
|
+
|
|
211
|
+
### Question:
|
|
212
|
+
{question}
|
|
213
|
+
|
|
214
|
+
### Response:
|
|
215
|
+
"""
|
|
216
|
+
|
|
188
217
|
QUERY_CONTEXTUALIZATION_PROMPT = """
|
|
189
218
|
# Role
|
|
190
219
|
You are a query rewriter for a Retrieval-Augmented Generation (RAG) system in high-energy physics (ATLAS/CERN).
|