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.
Files changed (48) hide show
  1. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/PKG-INFO +58 -1
  2. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/README.md +57 -0
  3. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/advanced.py +38 -25
  4. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/conversational_graph.py +202 -10
  5. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/groq.py +112 -3
  6. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/prompt/starters.py +29 -0
  7. chatlas_chains-0.3.0/chATLAS_Chains/router.py +209 -0
  8. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/PKG-INFO +58 -1
  9. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/SOURCES.txt +3 -0
  10. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/pyproject.toml +1 -1
  11. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_chat_model_kwargs.py +4 -2
  12. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_groq.py +23 -4
  13. chatlas_chains-0.3.0/tests/test_router.py +230 -0
  14. chatlas_chains-0.3.0/tests/test_router_litellm.py +40 -0
  15. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/LICENSE +0 -0
  16. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/__init__.py +0 -0
  17. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/__init__.py +0 -0
  18. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/basic.py +2 -2
  19. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/basic_graph.py +0 -0
  20. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
  21. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
  22. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/documents/rerank.py +0 -0
  23. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/documents/rrf.py +0 -0
  24. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/__init__.py +0 -0
  25. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/model_selection.py +0 -0
  26. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/llm/runnables.py +0 -0
  27. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/log.py +0 -0
  28. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/prompt/__init__.py +0 -0
  29. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  30. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/query/query_rewriting.py +0 -0
  31. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/search/__init__.py +0 -0
  32. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/search/basic.py +0 -0
  33. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/utils/__init__.py +0 -0
  34. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/utils/doc_utils.py +0 -0
  35. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains/vectorstore.py +0 -0
  36. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  37. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/requires.txt +0 -0
  38. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/chATLAS_Chains.egg-info/top_level.txt +0 -0
  39. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/setup.cfg +0 -0
  40. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/__init__.py +0 -0
  41. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/conftest.py +0 -0
  42. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_chains.py +0 -0
  43. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_conversational.py +0 -0
  44. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_llm_runnables.py +0 -0
  45. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_model_selection.py +0 -0
  46. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_rrf.py +0 -0
  47. {chatlas_chains-0.1.7 → chatlas_chains-0.3.0}/tests/test_search.py +0 -0
  48. {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.1.7
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
- graph.add_node("generate", generate_answer)
236
- # graph.add_node("assess", assess_answer)
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
- if not enable_rrf and not enable_reranking:
248
- # no document processing, go straight to generation
249
- graph.add_edge("retrieval", "generate")
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", "generate")
255
-
256
- elif not enable_rrf and enable_reranking:
257
- # 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"
258
254
  graph.add_edge("retrieval", "rerank")
259
- graph.add_edge("rerank", "generate")
260
-
261
255
  else:
262
- # retrieve → rrf → rerank → generate
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
- # generate at the end
268
- 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)
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 = query_rewriting_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
- contextualization_kwargs: ChatModelKwargs = {
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
- **contextualization_kwargs,
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
- standalone_question = str(response.content).strip()
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.outputs import ChatGeneration, ChatResult
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
- raise NotImplementedError("AccGPTChatGroqRateLimitAware does not support streaming yet.")
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).