chATLAS_Chains 0.3.1__tar.gz → 0.3.2__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.3.1 → chatlas_chains-0.3.2}/PKG-INFO +21 -11
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/README.md +20 -10
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/advanced.py +4 -4
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/conversational_graph.py +3 -3
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/llm/model_selection.py +16 -47
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/llm/runnables.py +12 -6
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/query/query_rewriting.py +2 -2
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/router.py +27 -11
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains.egg-info/PKG-INFO +21 -11
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/pyproject.toml +1 -1
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_chains.py +6 -6
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_conversational.py +15 -8
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_llm_runnables.py +3 -3
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_model_selection.py +19 -1
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_router.py +76 -11
- chatlas_chains-0.3.2/tests/test_router_litellm.py +101 -0
- chatlas_chains-0.3.1/tests/test_router_litellm.py +0 -42
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/LICENSE +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/basic.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/basic_graph.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/documents/rerank.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/documents/rrf.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/llm/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/llm/groq.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/log.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/prompt/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/prompt/starters.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/search/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/search/basic.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/utils/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/utils/doc_utils.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/vectorstore.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains.egg-info/SOURCES.txt +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains.egg-info/requires.txt +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains.egg-info/top_level.txt +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/setup.cfg +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/__init__.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/conftest.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_chat_model_kwargs.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_groq.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_rrf.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/tests/test_search.py +0 -0
- {chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/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.
|
|
3
|
+
Version: 0.3.2
|
|
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
|
|
@@ -66,9 +66,16 @@ 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
|
-
3. Using
|
|
69
|
+
3. Using chat models through [CERN AI Gateway](https://ml.docs.cern.ch/aigw/gettingstarted/).
|
|
70
|
+
The key variable retains its historical name so existing CI and OpenShift
|
|
71
|
+
configuration does not need to be renamed.
|
|
70
72
|
```bash
|
|
71
|
-
export CHATLAS_CHAINS_LITELLM_KEY="your
|
|
73
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your CERN AI Gateway key"
|
|
74
|
+
```
|
|
75
|
+
|
|
76
|
+
4. Optional proxy for LLM HTTP requests when direct endpoint access is unavailable.
|
|
77
|
+
```bash
|
|
78
|
+
export CHATLAS_PROXY="http://127.0.0.1:7890"
|
|
72
79
|
```
|
|
73
80
|
|
|
74
81
|
## Supported Chains
|
|
@@ -81,7 +88,7 @@ More details [here](chATLAS_Chains/chains/README.md)
|
|
|
81
88
|
## Prompt Routing
|
|
82
89
|
|
|
83
90
|
`chATLAS_Chains.router.route_user_prompt` classifies a raw user prompt before it enters the main RAG flow.
|
|
84
|
-
The router defaults to CERN
|
|
91
|
+
The router defaults to CERN AI Gateway and returns a validated `RouterDecision` with:
|
|
85
92
|
|
|
86
93
|
- `route`: one of `hep_atlas`, `person_lookup`, `dangerous`, `out_of_scope`,
|
|
87
94
|
or `unknown`
|
|
@@ -101,7 +108,7 @@ decision = route_user_prompt(
|
|
|
101
108
|
"Who is Jane Doe?",
|
|
102
109
|
chat_model_kwargs={
|
|
103
110
|
"service_provider": "litellm",
|
|
104
|
-
"proxy": "socks5h://localhost:1080", # optional when
|
|
111
|
+
"proxy": "socks5h://localhost:1080", # optional when direct access is unavailable
|
|
105
112
|
},
|
|
106
113
|
)
|
|
107
114
|
|
|
@@ -123,17 +130,20 @@ continue through the normal selected workflow.
|
|
|
123
130
|
|
|
124
131
|
Router context uses a hybrid policy: the current user prompt is authoritative,
|
|
125
132
|
and recent conversation context is only a bounded disambiguation aid for
|
|
126
|
-
follow-ups such as "What about Run 3?".
|
|
127
|
-
|
|
128
|
-
|
|
133
|
+
follow-ups such as "What about Run 3?". It keeps up to five previous
|
|
134
|
+
user/assistant turns, caps prior user messages at 600 characters and assistant
|
|
135
|
+
messages at 400 characters, drops assistant answers over 1,000 raw characters,
|
|
136
|
+
and caps retained history at 4,000 characters.
|
|
129
137
|
|
|
130
|
-
For live
|
|
138
|
+
For live AI Gateway calls, set the legacy-named credential:
|
|
131
139
|
|
|
132
140
|
```sh
|
|
133
|
-
export CHATLAS_CHAINS_LITELLM_KEY="your
|
|
141
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your CERN AI Gateway key"
|
|
134
142
|
```
|
|
135
143
|
|
|
136
|
-
|
|
144
|
+
Malformed model output is retried once with a strict JSON instruction. If the
|
|
145
|
+
gateway is unavailable or the retry is still invalid, the router raises an
|
|
146
|
+
error; the frontend logs it and fails open into the selected UI mode.
|
|
137
147
|
|
|
138
148
|
### Model Configuration in Chains
|
|
139
149
|
|
|
@@ -40,9 +40,16 @@ 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
|
-
3. Using
|
|
43
|
+
3. Using chat models through [CERN AI Gateway](https://ml.docs.cern.ch/aigw/gettingstarted/).
|
|
44
|
+
The key variable retains its historical name so existing CI and OpenShift
|
|
45
|
+
configuration does not need to be renamed.
|
|
44
46
|
```bash
|
|
45
|
-
export CHATLAS_CHAINS_LITELLM_KEY="your
|
|
47
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your CERN AI Gateway key"
|
|
48
|
+
```
|
|
49
|
+
|
|
50
|
+
4. Optional proxy for LLM HTTP requests when direct endpoint access is unavailable.
|
|
51
|
+
```bash
|
|
52
|
+
export CHATLAS_PROXY="http://127.0.0.1:7890"
|
|
46
53
|
```
|
|
47
54
|
|
|
48
55
|
## Supported Chains
|
|
@@ -55,7 +62,7 @@ More details [here](chATLAS_Chains/chains/README.md)
|
|
|
55
62
|
## Prompt Routing
|
|
56
63
|
|
|
57
64
|
`chATLAS_Chains.router.route_user_prompt` classifies a raw user prompt before it enters the main RAG flow.
|
|
58
|
-
The router defaults to CERN
|
|
65
|
+
The router defaults to CERN AI Gateway and returns a validated `RouterDecision` with:
|
|
59
66
|
|
|
60
67
|
- `route`: one of `hep_atlas`, `person_lookup`, `dangerous`, `out_of_scope`,
|
|
61
68
|
or `unknown`
|
|
@@ -75,7 +82,7 @@ decision = route_user_prompt(
|
|
|
75
82
|
"Who is Jane Doe?",
|
|
76
83
|
chat_model_kwargs={
|
|
77
84
|
"service_provider": "litellm",
|
|
78
|
-
"proxy": "socks5h://localhost:1080", # optional when
|
|
85
|
+
"proxy": "socks5h://localhost:1080", # optional when direct access is unavailable
|
|
79
86
|
},
|
|
80
87
|
)
|
|
81
88
|
|
|
@@ -97,17 +104,20 @@ continue through the normal selected workflow.
|
|
|
97
104
|
|
|
98
105
|
Router context uses a hybrid policy: the current user prompt is authoritative,
|
|
99
106
|
and recent conversation context is only a bounded disambiguation aid for
|
|
100
|
-
follow-ups such as "What about Run 3?".
|
|
101
|
-
|
|
102
|
-
|
|
107
|
+
follow-ups such as "What about Run 3?". It keeps up to five previous
|
|
108
|
+
user/assistant turns, caps prior user messages at 600 characters and assistant
|
|
109
|
+
messages at 400 characters, drops assistant answers over 1,000 raw characters,
|
|
110
|
+
and caps retained history at 4,000 characters.
|
|
103
111
|
|
|
104
|
-
For live
|
|
112
|
+
For live AI Gateway calls, set the legacy-named credential:
|
|
105
113
|
|
|
106
114
|
```sh
|
|
107
|
-
export CHATLAS_CHAINS_LITELLM_KEY="your
|
|
115
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your CERN AI Gateway key"
|
|
108
116
|
```
|
|
109
117
|
|
|
110
|
-
|
|
118
|
+
Malformed model output is retried once with a strict JSON instruction. If the
|
|
119
|
+
gateway is unavailable or the retry is still invalid, the router raises an
|
|
120
|
+
error; the frontend logs it and fails open into the selected UI mode.
|
|
111
121
|
|
|
112
122
|
### Model Configuration in Chains
|
|
113
123
|
|
|
@@ -24,7 +24,7 @@ from langgraph.graph.state import CompiledStateGraph
|
|
|
24
24
|
from chATLAS_Chains.documents.rerank import rerank_documents
|
|
25
25
|
from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
|
|
26
26
|
from chATLAS_Chains.llm.model_selection import (
|
|
27
|
-
|
|
27
|
+
DEFAULT_GROQ_MODEL,
|
|
28
28
|
ChatModelKwargs,
|
|
29
29
|
get_chat_model,
|
|
30
30
|
sanitize_chat_model_kwargs,
|
|
@@ -56,7 +56,7 @@ def advanced_rag(
|
|
|
56
56
|
enable_rrf: bool = False,
|
|
57
57
|
enable_reranking: bool = False,
|
|
58
58
|
# enable_self_evaluation: bool = False,
|
|
59
|
-
query_rewriting_model: str =
|
|
59
|
+
query_rewriting_model: str = DEFAULT_GROQ_MODEL,
|
|
60
60
|
query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
61
61
|
rerank_model: str = "cohere-rerank-3.5",
|
|
62
62
|
pinecone_api_key: str | None = None,
|
|
@@ -285,7 +285,7 @@ def stream_generate_answer(state: HybridGraphState, model, prompt_template: Chat
|
|
|
285
285
|
|
|
286
286
|
|
|
287
287
|
if __name__ == "__main__":
|
|
288
|
-
from chATLAS_Chains.llm.model_selection import
|
|
288
|
+
from chATLAS_Chains.llm.model_selection import DEFAULT_GROQ_MODEL
|
|
289
289
|
from chATLAS_Chains.vectorstore import get_vectorstore
|
|
290
290
|
|
|
291
291
|
twiki = get_vectorstore("twiki_prod")
|
|
@@ -294,7 +294,7 @@ if __name__ == "__main__":
|
|
|
294
294
|
# Create the hybrid graph
|
|
295
295
|
graph = advanced_rag(
|
|
296
296
|
vectorstore=[twiki, mkdocs],
|
|
297
|
-
model_name=
|
|
297
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
298
298
|
enable_query_rewriting=True,
|
|
299
299
|
enable_rrf=True,
|
|
300
300
|
enable_reranking=False,
|
|
@@ -18,7 +18,7 @@ from sentence_transformers.util import cos_sim
|
|
|
18
18
|
|
|
19
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
|
+
DEFAULT_GROQ_MODEL,
|
|
22
22
|
ChatModelKwargs,
|
|
23
23
|
get_chat_model,
|
|
24
24
|
sanitize_chat_model_kwargs,
|
|
@@ -134,7 +134,7 @@ def conversational_retrieval_graph(
|
|
|
134
134
|
enable_rrf: bool = True,
|
|
135
135
|
enable_reranking: bool = False,
|
|
136
136
|
enable_query_contextualization: bool = True,
|
|
137
|
-
query_rewriting_model: str =
|
|
137
|
+
query_rewriting_model: str = DEFAULT_GROQ_MODEL,
|
|
138
138
|
query_rewriting_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
139
139
|
contextualization_model: str | None = None,
|
|
140
140
|
contextualization_chat_model_kwargs: ChatModelKwargs | None = None,
|
|
@@ -528,7 +528,7 @@ if __name__ == "__main__":
|
|
|
528
528
|
|
|
529
529
|
graph = conversational_retrieval_graph(
|
|
530
530
|
vectorstore=vectorstore,
|
|
531
|
-
model_name=
|
|
531
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
532
532
|
checkpointer=checkpointer,
|
|
533
533
|
enable_query_rewriting=False,
|
|
534
534
|
enable_rrf=True,
|
|
@@ -6,42 +6,13 @@ from chATLAS_Chains.llm.runnables import groq_runnable, litellm_runnable, openai
|
|
|
6
6
|
|
|
7
7
|
logger = logging.getLogger(__name__)
|
|
8
8
|
|
|
9
|
-
#
|
|
9
|
+
# CERN AI Gateway chat models. The constant retains its historical name for
|
|
10
|
+
# compatibility with callers that import it directly.
|
|
11
|
+
# See https://ml.docs.cern.ch/aigw/models/
|
|
10
12
|
LITELLM_MODELS = [
|
|
11
|
-
"
|
|
12
|
-
"hf-
|
|
13
|
+
"qwen3.8-27b-fp16",
|
|
14
|
+
"hf-qwen3-32b-awq",
|
|
13
15
|
"gpt-oss-20b",
|
|
14
|
-
"gpt-4",
|
|
15
|
-
"gpt-4-turbo",
|
|
16
|
-
"gpt-4.1-nano",
|
|
17
|
-
"gpt-5",
|
|
18
|
-
"gpt-5.1",
|
|
19
|
-
"gpt-5.2",
|
|
20
|
-
"gpt-4.1",
|
|
21
|
-
"e5-large-v2",
|
|
22
|
-
"accgpt",
|
|
23
|
-
"magistral-medium-latest",
|
|
24
|
-
"codestral-latest",
|
|
25
|
-
"devstral-medium-latest",
|
|
26
|
-
"devstral-small-latest",
|
|
27
|
-
"magistral-small-latest",
|
|
28
|
-
"ministral-3b-latest",
|
|
29
|
-
"ministral-8b-latest",
|
|
30
|
-
"mistral-large-latest",
|
|
31
|
-
"mistral-medium-latest",
|
|
32
|
-
"mistral-small-latest",
|
|
33
|
-
"mistral-tiny-latest",
|
|
34
|
-
"mistral-moderation-latest",
|
|
35
|
-
"mistral-ocr-latest",
|
|
36
|
-
"pixtral-12b-latest",
|
|
37
|
-
"pixtral-large-latest",
|
|
38
|
-
"voxtral-mini-latest",
|
|
39
|
-
"voxtral-small-latest",
|
|
40
|
-
"llama-3.1-8b-instant",
|
|
41
|
-
"openai/gpt-oss-120b",
|
|
42
|
-
"moonshotai/kimi-k2-instruct-0905",
|
|
43
|
-
"meta-llama/llama-4-scout-17b-16e-instruct",
|
|
44
|
-
"meta-llama/llama-4-maverick-17b-128e-instruct",
|
|
45
16
|
]
|
|
46
17
|
|
|
47
18
|
# limit what OpenAI models we support (to control costs)
|
|
@@ -59,11 +30,9 @@ OPENAI_MODELS = [
|
|
|
59
30
|
]
|
|
60
31
|
|
|
61
32
|
# see Groq docs: https://console.groq.com/docs/models
|
|
33
|
+
DEFAULT_GROQ_MODEL = "openai/gpt-oss-120b"
|
|
62
34
|
GROQ_PRODUCTION_MODELS = [
|
|
63
|
-
|
|
64
|
-
"llama-3.3-70b-versatile",
|
|
65
|
-
"meta-llama/llama-guard-4-12b",
|
|
66
|
-
"openai/gpt-oss-120b",
|
|
35
|
+
DEFAULT_GROQ_MODEL,
|
|
67
36
|
"openai/gpt-oss-20b",
|
|
68
37
|
]
|
|
69
38
|
|
|
@@ -172,16 +141,16 @@ def get_chat_model(
|
|
|
172
141
|
model_name: Model to load
|
|
173
142
|
|
|
174
143
|
service_provider: Which provider to use for this model.
|
|
175
|
-
Used when a model is available from multiple providers
|
|
144
|
+
Used when a model is available from multiple providers.
|
|
176
145
|
|
|
177
146
|
api_key: Optional API key for the provider. If not provided, will look for environment variables:
|
|
178
147
|
- OpenAI: `CHATLAS_OPENAI_KEY`
|
|
179
148
|
- Groq: `CHATLAS_GROQ_KEY`
|
|
180
|
-
-
|
|
149
|
+
- CERN AI Gateway: `CHATLAS_CHAINS_LITELLM_KEY` (legacy name)
|
|
181
150
|
|
|
182
|
-
base_url: Optional base URL for the provider. Only needed for Groq and
|
|
151
|
+
base_url: Optional base URL for the provider. Only needed for Groq and CERN AI Gateway
|
|
183
152
|
- Groq: `CHATLAS_GROQ_BASE_URL`
|
|
184
|
-
-
|
|
153
|
+
- CERN AI Gateway: defaults to "https://aigw.cern.ch/v1" if not provided
|
|
185
154
|
|
|
186
155
|
max_tokens: Maximum number of tokens to generate. If ``None``, uses the
|
|
187
156
|
model default (or provider default behavior).
|
|
@@ -212,7 +181,8 @@ def get_chat_model(
|
|
|
212
181
|
# warn if also avaiable from elsewhere
|
|
213
182
|
if in_openai or in_litellm:
|
|
214
183
|
print(
|
|
215
|
-
f"Model '{model_name}' is also available from
|
|
184
|
+
f"Model '{model_name}' is also available from CERN AI Gateway or OpenAI. Defaulting to Groq. "
|
|
185
|
+
"To select another provider, pass service_provider explicitly."
|
|
216
186
|
)
|
|
217
187
|
|
|
218
188
|
elif in_litellm:
|
|
@@ -251,7 +221,7 @@ def get_chat_model(
|
|
|
251
221
|
)
|
|
252
222
|
|
|
253
223
|
elif service_provider == "litellm":
|
|
254
|
-
runnable_base_url = base_url if base_url is not None else "https://
|
|
224
|
+
runnable_base_url = base_url if base_url is not None else "https://aigw.cern.ch/v1"
|
|
255
225
|
llm = litellm_runnable(
|
|
256
226
|
model_name,
|
|
257
227
|
base_url=runnable_base_url,
|
|
@@ -277,16 +247,15 @@ if __name__ == "__main__":
|
|
|
277
247
|
|
|
278
248
|
service_provider_model = [
|
|
279
249
|
# ("", "llama-3.1-8b-instruct"),
|
|
280
|
-
# ("", "llama-3.1-8b-instant"),
|
|
281
250
|
# ("", "gpt-oss-20b"),
|
|
282
251
|
# ("openai", "gpt-5-mini"),
|
|
283
|
-
# ("groq",
|
|
252
|
+
# ("groq", DEFAULT_GROQ_MODEL),
|
|
284
253
|
("litellm", "gpt-oss-20b"),
|
|
285
254
|
]
|
|
286
255
|
|
|
287
256
|
for service_provider, model_name in service_provider_model:
|
|
288
257
|
if service_provider == "litellm":
|
|
289
|
-
print(f"Testing
|
|
258
|
+
print(f"Testing CERN AI Gateway model '{model_name}' with proxy...")
|
|
290
259
|
llm = get_chat_model(model_name, service_provider, proxy="socks5h://localhost:1080")
|
|
291
260
|
else:
|
|
292
261
|
print(f"Testing model '{model_name}' from provider '{service_provider or 'auto-detected'}'...")
|
|
@@ -25,7 +25,7 @@ def llm_runnable(
|
|
|
25
25
|
Base function to return a LangChain runnable. Wrapped by the various LLM providers.
|
|
26
26
|
|
|
27
27
|
Args:
|
|
28
|
-
proxy: Optional proxy URL to use for requests
|
|
28
|
+
proxy: Optional proxy URL to use for requests when direct access is unavailable.
|
|
29
29
|
"""
|
|
30
30
|
|
|
31
31
|
# Build an httpx client with explicit proxy behavior
|
|
@@ -106,19 +106,24 @@ def groq_runnable(
|
|
|
106
106
|
|
|
107
107
|
def litellm_runnable(
|
|
108
108
|
model_name: str,
|
|
109
|
-
base_url: str = "https://
|
|
109
|
+
base_url: str = "https://aigw.cern.ch/v1",
|
|
110
110
|
api_key: str | None = None,
|
|
111
111
|
temperature: float | None = None,
|
|
112
112
|
max_tokens: int | None = None,
|
|
113
113
|
proxy: str | None = None,
|
|
114
114
|
):
|
|
115
115
|
"""
|
|
116
|
-
Returns a LangChain runnable using
|
|
116
|
+
Returns a LangChain runnable using CERN AI Gateway.
|
|
117
|
+
|
|
118
|
+
The function and environment-variable names retain ``litellm`` for
|
|
119
|
+
deployment compatibility. AI Gateway exposes the same OpenAI-compatible
|
|
120
|
+
protocol, so existing callers do not need a credential migration.
|
|
117
121
|
|
|
118
122
|
Args:
|
|
119
123
|
model_name: Model to use
|
|
120
|
-
base_url:
|
|
121
|
-
api_key:
|
|
124
|
+
base_url: AI Gateway's OpenAI-compatible API URL.
|
|
125
|
+
api_key: AI Gateway API key. If None, reads the legacy
|
|
126
|
+
CHATLAS_CHAINS_LITELLM_KEY environment variable.
|
|
122
127
|
temperature: Model temperature. Defaults to None.
|
|
123
128
|
max_tokens: Maximum number of tokens to generate. If ``None``, uses the
|
|
124
129
|
model default (or provider default behavior).
|
|
@@ -129,7 +134,8 @@ def litellm_runnable(
|
|
|
129
134
|
api_key = os.getenv("CHATLAS_CHAINS_LITELLM_KEY")
|
|
130
135
|
if not api_key or not api_key.strip():
|
|
131
136
|
raise ValueError(
|
|
132
|
-
"Missing API key
|
|
137
|
+
"Missing CERN AI Gateway API key. Set the CHATLAS_CHAINS_LITELLM_KEY environment variable "
|
|
138
|
+
"or provide it as an argument."
|
|
133
139
|
)
|
|
134
140
|
|
|
135
141
|
return llm_runnable(model_name, base_url, api_key, temperature, max_tokens=max_tokens, proxy=proxy)
|
|
@@ -45,9 +45,9 @@ def rewrite_query(question: str, model: BaseChatModel, rewrite_prompt: str = QUE
|
|
|
45
45
|
|
|
46
46
|
if __name__ == "__main__":
|
|
47
47
|
# Example usage
|
|
48
|
-
from chATLAS_Chains.llm.model_selection import
|
|
48
|
+
from chATLAS_Chains.llm.model_selection import DEFAULT_GROQ_MODEL, get_chat_model
|
|
49
49
|
|
|
50
|
-
llm = get_chat_model(model_name=
|
|
50
|
+
llm = get_chat_model(model_name=DEFAULT_GROQ_MODEL)
|
|
51
51
|
|
|
52
52
|
question = "What is the Drell-Yan dilepton cross section?"
|
|
53
53
|
print(f"Original question: {question}")
|
|
@@ -1,6 +1,7 @@
|
|
|
1
|
-
"""
|
|
1
|
+
"""CERN AI Gateway intent router for chATLAS requests."""
|
|
2
2
|
|
|
3
3
|
import json
|
|
4
|
+
import logging
|
|
4
5
|
import re
|
|
5
6
|
from collections.abc import Sequence
|
|
6
7
|
from typing import Any, Literal, TypedDict
|
|
@@ -17,18 +18,25 @@ RouterRoute = Literal[
|
|
|
17
18
|
"unknown",
|
|
18
19
|
]
|
|
19
20
|
|
|
20
|
-
|
|
21
|
+
logger = logging.getLogger(__name__)
|
|
22
|
+
|
|
23
|
+
ROUTER_PROMPT_VERSION = "2026-09-18.1"
|
|
21
24
|
_DEFAULT_ROUTER_MODEL = "gpt-oss-20b"
|
|
22
25
|
_DEFAULT_ROUTER_KWARGS: ChatModelKwargs = {
|
|
23
26
|
"service_provider": "litellm",
|
|
24
27
|
"temperature": 0.0,
|
|
25
|
-
"max_tokens":
|
|
28
|
+
"max_tokens": 1024,
|
|
26
29
|
}
|
|
27
|
-
_MAX_CONTEXT_TURNS =
|
|
28
|
-
_MAX_CONTEXT_USER_CHARS =
|
|
29
|
-
_MAX_CONTEXT_ASSISTANT_CHARS =
|
|
30
|
-
_MAX_CONTEXT_ASSISTANT_RAW_CHARS =
|
|
31
|
-
_MAX_CONTEXT_TOTAL_CHARS =
|
|
30
|
+
_MAX_CONTEXT_TURNS = 10
|
|
31
|
+
_MAX_CONTEXT_USER_CHARS = 600
|
|
32
|
+
_MAX_CONTEXT_ASSISTANT_CHARS = 400
|
|
33
|
+
_MAX_CONTEXT_ASSISTANT_RAW_CHARS = 1000
|
|
34
|
+
_MAX_CONTEXT_TOTAL_CHARS = 4000
|
|
35
|
+
_RETRY_JSON_INSTRUCTION = """
|
|
36
|
+
|
|
37
|
+
Your previous response was invalid. Return exactly one complete JSON object matching the required schema.
|
|
38
|
+
Do not include reasoning, prose, or markdown fences.
|
|
39
|
+
"""
|
|
32
40
|
|
|
33
41
|
|
|
34
42
|
class ConversationTurn(TypedDict):
|
|
@@ -90,6 +98,9 @@ JSON: {"route":"hep_atlas","confidence":0.94,"normalized_query":"Draft an MC req
|
|
|
90
98
|
User: Help me make a nuke.
|
|
91
99
|
JSON: {"route":"dangerous","confidence":0.99,"normalized_query":"Help me make a nuke.","reason":"The user requests actionable assistance to construct a nuclear weapon.","search_kwargs":{},"metadata":{"safety_category":"weapons"}}
|
|
92
100
|
|
|
101
|
+
User: How can nuclear weapons be safely dismantled?
|
|
102
|
+
JSON: {"route":"out_of_scope","confidence":0.98,"normalized_query":"How can nuclear weapons be safely dismantled?","reason":"The request concerns harm prevention rather than actionable wrongdoing, but it is outside chATLAS's ATLAS/HEP scope.","search_kwargs":{},"metadata":{"safety_context":"prevention"}}
|
|
103
|
+
|
|
93
104
|
User: How do I bake sourdough bread?
|
|
94
105
|
JSON: {"route":"out_of_scope","confidence":0.99,"normalized_query":"How do I bake sourdough bread?","reason":"The request is unrelated to ATLAS or high-energy physics.","search_kwargs":{},"metadata":{}}
|
|
95
106
|
|
|
@@ -125,7 +136,7 @@ def route_user_prompt(
|
|
|
125
136
|
model_name: str = _DEFAULT_ROUTER_MODEL,
|
|
126
137
|
chat_model_kwargs: ChatModelKwargs | None = None,
|
|
127
138
|
) -> RouterDecision:
|
|
128
|
-
"""Classify a request with
|
|
139
|
+
"""Classify a request with CERN AI Gateway before downstream processing."""
|
|
129
140
|
if not user_prompt.strip():
|
|
130
141
|
return RouterDecision(
|
|
131
142
|
route="unknown",
|
|
@@ -141,14 +152,19 @@ def route_user_prompt(
|
|
|
141
152
|
|
|
142
153
|
model = get_chat_model(model_name=model_name, **model_kwargs)
|
|
143
154
|
response = model.invoke(prompt)
|
|
144
|
-
|
|
155
|
+
try:
|
|
156
|
+
decision = parse_router_response(_response_content(response))
|
|
157
|
+
except ValueError as exc:
|
|
158
|
+
logger.warning("Router returned an invalid response; retrying once: %s", exc)
|
|
159
|
+
retry_response = model.invoke(prompt + _RETRY_JSON_INSTRUCTION)
|
|
160
|
+
decision = parse_router_response(_response_content(retry_response))
|
|
145
161
|
if not decision.normalized_query:
|
|
146
162
|
decision.normalized_query = user_prompt.strip()
|
|
147
163
|
return decision
|
|
148
164
|
|
|
149
165
|
|
|
150
166
|
def parse_router_response(response_text: str) -> RouterDecision:
|
|
151
|
-
"""Parse and validate
|
|
167
|
+
"""Parse and validate an AI Gateway router response."""
|
|
152
168
|
payload = _extract_json_object(response_text)
|
|
153
169
|
payload = _normalize_disabled_routes(json.loads(payload))
|
|
154
170
|
return RouterDecision.model_validate(payload)
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: chATLAS_Chains
|
|
3
|
-
Version: 0.3.
|
|
3
|
+
Version: 0.3.2
|
|
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
|
|
@@ -66,9 +66,16 @@ 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
|
-
3. Using
|
|
69
|
+
3. Using chat models through [CERN AI Gateway](https://ml.docs.cern.ch/aigw/gettingstarted/).
|
|
70
|
+
The key variable retains its historical name so existing CI and OpenShift
|
|
71
|
+
configuration does not need to be renamed.
|
|
70
72
|
```bash
|
|
71
|
-
export CHATLAS_CHAINS_LITELLM_KEY="your
|
|
73
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your CERN AI Gateway key"
|
|
74
|
+
```
|
|
75
|
+
|
|
76
|
+
4. Optional proxy for LLM HTTP requests when direct endpoint access is unavailable.
|
|
77
|
+
```bash
|
|
78
|
+
export CHATLAS_PROXY="http://127.0.0.1:7890"
|
|
72
79
|
```
|
|
73
80
|
|
|
74
81
|
## Supported Chains
|
|
@@ -81,7 +88,7 @@ More details [here](chATLAS_Chains/chains/README.md)
|
|
|
81
88
|
## Prompt Routing
|
|
82
89
|
|
|
83
90
|
`chATLAS_Chains.router.route_user_prompt` classifies a raw user prompt before it enters the main RAG flow.
|
|
84
|
-
The router defaults to CERN
|
|
91
|
+
The router defaults to CERN AI Gateway and returns a validated `RouterDecision` with:
|
|
85
92
|
|
|
86
93
|
- `route`: one of `hep_atlas`, `person_lookup`, `dangerous`, `out_of_scope`,
|
|
87
94
|
or `unknown`
|
|
@@ -101,7 +108,7 @@ decision = route_user_prompt(
|
|
|
101
108
|
"Who is Jane Doe?",
|
|
102
109
|
chat_model_kwargs={
|
|
103
110
|
"service_provider": "litellm",
|
|
104
|
-
"proxy": "socks5h://localhost:1080", # optional when
|
|
111
|
+
"proxy": "socks5h://localhost:1080", # optional when direct access is unavailable
|
|
105
112
|
},
|
|
106
113
|
)
|
|
107
114
|
|
|
@@ -123,17 +130,20 @@ continue through the normal selected workflow.
|
|
|
123
130
|
|
|
124
131
|
Router context uses a hybrid policy: the current user prompt is authoritative,
|
|
125
132
|
and recent conversation context is only a bounded disambiguation aid for
|
|
126
|
-
follow-ups such as "What about Run 3?".
|
|
127
|
-
|
|
128
|
-
|
|
133
|
+
follow-ups such as "What about Run 3?". It keeps up to five previous
|
|
134
|
+
user/assistant turns, caps prior user messages at 600 characters and assistant
|
|
135
|
+
messages at 400 characters, drops assistant answers over 1,000 raw characters,
|
|
136
|
+
and caps retained history at 4,000 characters.
|
|
129
137
|
|
|
130
|
-
For live
|
|
138
|
+
For live AI Gateway calls, set the legacy-named credential:
|
|
131
139
|
|
|
132
140
|
```sh
|
|
133
|
-
export CHATLAS_CHAINS_LITELLM_KEY="your
|
|
141
|
+
export CHATLAS_CHAINS_LITELLM_KEY="your CERN AI Gateway key"
|
|
134
142
|
```
|
|
135
143
|
|
|
136
|
-
|
|
144
|
+
Malformed model output is retried once with a strict JSON instruction. If the
|
|
145
|
+
gateway is unavailable or the retry is still invalid, the router raises an
|
|
146
|
+
error; the frontend logs it and fails open into the selected UI mode.
|
|
137
147
|
|
|
138
148
|
### Model Configuration in Chains
|
|
139
149
|
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "chATLAS_Chains"
|
|
7
|
-
version = "0.3.
|
|
7
|
+
version = "0.3.2"
|
|
8
8
|
description = "A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project."
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.11"
|
|
@@ -13,7 +13,7 @@ from pydantic import ValidationError
|
|
|
13
13
|
|
|
14
14
|
from chATLAS_Chains.chains.advanced import advanced_rag
|
|
15
15
|
from chATLAS_Chains.chains.basic import basic_retrieval_chain
|
|
16
|
-
from chATLAS_Chains.llm.model_selection import
|
|
16
|
+
from chATLAS_Chains.llm.model_selection import DEFAULT_GROQ_MODEL
|
|
17
17
|
from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
|
|
18
18
|
|
|
19
19
|
GROQ_MODEL_KWARGS = {"service_provider": "groq"}
|
|
@@ -44,7 +44,7 @@ def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore):
|
|
|
44
44
|
chain = basic_retrieval_chain(
|
|
45
45
|
prompt=CHAT_PROMPT_TEMPLATE,
|
|
46
46
|
vectorstore=twiki_vectorstore,
|
|
47
|
-
model_name=
|
|
47
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
48
48
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
49
49
|
)
|
|
50
50
|
output = chain.invoke("What is the Higgs boson?")
|
|
@@ -60,7 +60,7 @@ def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores):
|
|
|
60
60
|
chain = basic_retrieval_chain(
|
|
61
61
|
prompt=CHAT_PROMPT_TEMPLATE,
|
|
62
62
|
vectorstore=three_vectorstores,
|
|
63
|
-
model_name=
|
|
63
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
64
64
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
65
65
|
)
|
|
66
66
|
output = chain.invoke("What is the Higgs boson?")
|
|
@@ -90,7 +90,7 @@ def test_advanced_rag_chain(twiki_vectorstore):
|
|
|
90
90
|
chain = advanced_rag(
|
|
91
91
|
prompt=CHAT_PROMPT_TEMPLATE,
|
|
92
92
|
vectorstore=twiki_vectorstore,
|
|
93
|
-
model_name=
|
|
93
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
94
94
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
95
95
|
)
|
|
96
96
|
output = chain.invoke(
|
|
@@ -109,7 +109,7 @@ def test_advanced_rag_chain(twiki_vectorstore):
|
|
|
109
109
|
chain = advanced_rag(
|
|
110
110
|
prompt=CHAT_PROMPT_TEMPLATE,
|
|
111
111
|
vectorstore=twiki_vectorstore,
|
|
112
|
-
model_name=
|
|
112
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
113
113
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
114
114
|
enable_query_rewriting=True,
|
|
115
115
|
query_rewriting_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
@@ -141,7 +141,7 @@ def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores):
|
|
|
141
141
|
chain = advanced_rag(
|
|
142
142
|
prompt=CHAT_PROMPT_TEMPLATE,
|
|
143
143
|
vectorstore=three_vectorstores,
|
|
144
|
-
model_name=
|
|
144
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
145
145
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
146
146
|
)
|
|
147
147
|
output = chain.invoke(
|
|
@@ -7,7 +7,7 @@ from langchain_core.messages import HumanMessage
|
|
|
7
7
|
from langgraph.checkpoint.memory import MemorySaver
|
|
8
8
|
|
|
9
9
|
from chATLAS_Chains.chains.conversational_graph import conversational_retrieval_graph
|
|
10
|
-
from chATLAS_Chains.llm.model_selection import
|
|
10
|
+
from chATLAS_Chains.llm.model_selection import DEFAULT_GROQ_MODEL
|
|
11
11
|
|
|
12
12
|
GROQ_MODEL_KWARGS = {"service_provider": "groq"}
|
|
13
13
|
|
|
@@ -28,9 +28,10 @@ def test_pronoun_resolution(twiki_vectorstore, checkpointer, search_kwargs):
|
|
|
28
28
|
"""Test pronoun 'them' resolves to TWiki rules from previous turn."""
|
|
29
29
|
graph = conversational_retrieval_graph(
|
|
30
30
|
vectorstore=twiki_vectorstore,
|
|
31
|
-
model_name=
|
|
31
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
32
32
|
checkpointer=checkpointer,
|
|
33
33
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
34
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
34
35
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
35
36
|
)
|
|
36
37
|
|
|
@@ -58,9 +59,10 @@ def test_incomplete_reference(twiki_vectorstore, checkpointer, search_kwargs):
|
|
|
58
59
|
"""Test ordinal reference 'the third one' resolves to specific list item."""
|
|
59
60
|
graph = conversational_retrieval_graph(
|
|
60
61
|
vectorstore=twiki_vectorstore,
|
|
61
|
-
model_name=
|
|
62
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
62
63
|
checkpointer=checkpointer,
|
|
63
64
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
65
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
64
66
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
65
67
|
)
|
|
66
68
|
|
|
@@ -91,9 +93,10 @@ def test_clarification_refinement(twiki_vectorstore, checkpointer, search_kwargs
|
|
|
91
93
|
"""Test clarification 'Using WikiWords' refines the previous question."""
|
|
92
94
|
graph = conversational_retrieval_graph(
|
|
93
95
|
vectorstore=twiki_vectorstore,
|
|
94
|
-
model_name=
|
|
96
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
95
97
|
checkpointer=checkpointer,
|
|
96
98
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
99
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
97
100
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
98
101
|
)
|
|
99
102
|
|
|
@@ -118,9 +121,10 @@ def test_topic_shift(twiki_vectorstore, checkpointer, search_kwargs):
|
|
|
118
121
|
"""Test topic shift - new topic should not include previous context."""
|
|
119
122
|
graph = conversational_retrieval_graph(
|
|
120
123
|
vectorstore=twiki_vectorstore,
|
|
121
|
-
model_name=
|
|
124
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
122
125
|
checkpointer=checkpointer,
|
|
123
126
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
127
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
124
128
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
125
129
|
)
|
|
126
130
|
|
|
@@ -145,9 +149,10 @@ def test_generalization(twiki_vectorstore, checkpointer, search_kwargs):
|
|
|
145
149
|
"""Test generalization - broader question should not inherit specific context."""
|
|
146
150
|
graph = conversational_retrieval_graph(
|
|
147
151
|
vectorstore=twiki_vectorstore,
|
|
148
|
-
model_name=
|
|
152
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
149
153
|
checkpointer=checkpointer,
|
|
150
154
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
155
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
151
156
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
152
157
|
)
|
|
153
158
|
|
|
@@ -175,9 +180,10 @@ def test_spelling_and_grammar_correction(twiki_vectorstore, checkpointer, search
|
|
|
175
180
|
"""Test spelling and grammar errors are corrected while preserving abbreviations."""
|
|
176
181
|
graph = conversational_retrieval_graph(
|
|
177
182
|
vectorstore=twiki_vectorstore,
|
|
178
|
-
model_name=
|
|
183
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
179
184
|
checkpointer=checkpointer,
|
|
180
185
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
186
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
181
187
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
182
188
|
)
|
|
183
189
|
|
|
@@ -205,9 +211,10 @@ def test_hard_turn_limit(twiki_vectorstore, checkpointer, search_kwargs):
|
|
|
205
211
|
# Set hard limit to 2 turns
|
|
206
212
|
graph = conversational_retrieval_graph(
|
|
207
213
|
vectorstore=twiki_vectorstore,
|
|
208
|
-
model_name=
|
|
214
|
+
model_name=DEFAULT_GROQ_MODEL,
|
|
209
215
|
checkpointer=checkpointer,
|
|
210
216
|
chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
217
|
+
contextualization_model=DEFAULT_GROQ_MODEL,
|
|
211
218
|
contextualization_chat_model_kwargs=GROQ_MODEL_KWARGS,
|
|
212
219
|
max_turns=2,
|
|
213
220
|
)
|
|
@@ -32,11 +32,11 @@ def test_litellm_model_uses_env_var_when_api_key_not_provided(monkeypatch):
|
|
|
32
32
|
monkeypatch.setenv("CHATLAS_CHAINS_LITELLM_KEY", " sk-env-key ")
|
|
33
33
|
|
|
34
34
|
llm = litellm_runnable(
|
|
35
|
-
model_name="hf-
|
|
35
|
+
model_name="hf-qwen3-32b-awq",
|
|
36
36
|
base_url="https://example.com/v1",
|
|
37
37
|
)
|
|
38
38
|
|
|
39
|
-
assert llm.model_name == "hf-
|
|
39
|
+
assert llm.model_name == "hf-qwen3-32b-awq"
|
|
40
40
|
assert _secret_to_str(llm.openai_api_key) == "sk-env-key"
|
|
41
41
|
assert llm.openai_api_base == "https://example.com/v1"
|
|
42
42
|
|
|
@@ -44,5 +44,5 @@ def test_litellm_model_uses_env_var_when_api_key_not_provided(monkeypatch):
|
|
|
44
44
|
def test_litellm_model_raises_if_no_key_in_arg_or_env(monkeypatch):
|
|
45
45
|
monkeypatch.delenv("CHATLAS_CHAINS_LITELLM_KEY", raising=False)
|
|
46
46
|
|
|
47
|
-
with pytest.raises(ValueError, match="Missing API key
|
|
47
|
+
with pytest.raises(ValueError, match="Missing CERN AI Gateway API key"):
|
|
48
48
|
litellm_runnable(model_name="gpt-4")
|
|
@@ -100,7 +100,7 @@ def test_litellm_called_with_base_url(monkeypatch):
|
|
|
100
100
|
|
|
101
101
|
def fake_litellm_runnable(
|
|
102
102
|
model_name_arg,
|
|
103
|
-
base_url="https://
|
|
103
|
+
base_url="https://aigw.cern.ch/v1",
|
|
104
104
|
api_key=None,
|
|
105
105
|
temperature=None,
|
|
106
106
|
max_tokens=None,
|
|
@@ -127,3 +127,21 @@ def test_litellm_called_with_base_url(monkeypatch):
|
|
|
127
127
|
assert llm is not None
|
|
128
128
|
assert called["base_url"] == "https://litellm.test"
|
|
129
129
|
assert called["temperature"] == 0.3
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def test_litellm_defaults_to_cern_ai_gateway(monkeypatch):
|
|
133
|
+
called = {}
|
|
134
|
+
|
|
135
|
+
def fake_litellm_runnable(model_name_arg, base_url, **_kwargs):
|
|
136
|
+
called["model_name"] = model_name_arg
|
|
137
|
+
called["base_url"] = base_url
|
|
138
|
+
return object()
|
|
139
|
+
|
|
140
|
+
monkeypatch.setattr(ms, "litellm_runnable", fake_litellm_runnable)
|
|
141
|
+
|
|
142
|
+
ms.get_chat_model("gpt-oss-20b", service_provider="litellm", api_key="k")
|
|
143
|
+
|
|
144
|
+
assert called == {
|
|
145
|
+
"model_name": "gpt-oss-20b",
|
|
146
|
+
"base_url": "https://aigw.cern.ch/v1",
|
|
147
|
+
}
|
|
@@ -7,13 +7,17 @@ from chATLAS_Chains.router import RouterDecision, parse_router_response, route_u
|
|
|
7
7
|
|
|
8
8
|
|
|
9
9
|
class FakeRouterModel:
|
|
10
|
-
def __init__(self, content: str):
|
|
11
|
-
self.
|
|
10
|
+
def __init__(self, content: str | Exception | list[str | Exception]):
|
|
11
|
+
self.contents = content if isinstance(content, list) else [content]
|
|
12
12
|
self.prompts: list[str] = []
|
|
13
13
|
|
|
14
14
|
def invoke(self, prompt: str):
|
|
15
15
|
self.prompts.append(prompt)
|
|
16
|
-
|
|
16
|
+
index = min(len(self.prompts) - 1, len(self.contents) - 1)
|
|
17
|
+
content = self.contents[index]
|
|
18
|
+
if isinstance(content, Exception):
|
|
19
|
+
raise content
|
|
20
|
+
return SimpleNamespace(content=content)
|
|
17
21
|
|
|
18
22
|
|
|
19
23
|
def test_route_user_prompt_uses_litellm_defaults(monkeypatch):
|
|
@@ -41,7 +45,7 @@ def test_route_user_prompt_uses_litellm_defaults(monkeypatch):
|
|
|
41
45
|
)
|
|
42
46
|
assert captured == {
|
|
43
47
|
"model_name": "gpt-oss-20b",
|
|
44
|
-
"kwargs": {"service_provider": "litellm", "temperature": 0.0, "max_tokens":
|
|
48
|
+
"kwargs": {"service_provider": "litellm", "temperature": 0.0, "max_tokens": 1024},
|
|
45
49
|
}
|
|
46
50
|
assert "Return JSON only" in fake_model.prompts[0]
|
|
47
51
|
assert "selected UI mode" not in fake_model.prompts[0]
|
|
@@ -62,7 +66,7 @@ def test_route_user_prompt_allows_litellm_overrides(monkeypatch):
|
|
|
62
66
|
monkeypatch.setattr(router, "get_chat_model", fake_get_chat_model)
|
|
63
67
|
decision = route_user_prompt(
|
|
64
68
|
"Who is Jane Doe?",
|
|
65
|
-
model_name="hf-
|
|
69
|
+
model_name="hf-qwen3-32b-awq",
|
|
66
70
|
chat_model_kwargs={
|
|
67
71
|
"api_key": "sk-test",
|
|
68
72
|
"base_url": "https://litellm.example/v1",
|
|
@@ -76,7 +80,7 @@ def test_route_user_prompt_allows_litellm_overrides(monkeypatch):
|
|
|
76
80
|
assert captured["kwargs"] == {
|
|
77
81
|
"service_provider": "litellm",
|
|
78
82
|
"temperature": 0.2,
|
|
79
|
-
"max_tokens":
|
|
83
|
+
"max_tokens": 1024,
|
|
80
84
|
"api_key": "sk-test",
|
|
81
85
|
"base_url": "https://litellm.example/v1",
|
|
82
86
|
"proxy": "socks5h://localhost:1080",
|
|
@@ -127,7 +131,7 @@ def test_route_user_prompt_uses_hybrid_context_policy(monkeypatch):
|
|
|
127
131
|
assert "How do I apply a GRL?" in prompt
|
|
128
132
|
assert "Use the GoodRunsListSelectionTool" not in prompt
|
|
129
133
|
assert "Use the JetEtmiss recommendations." in prompt
|
|
130
|
-
assert len(prompt) <
|
|
134
|
+
assert len(prompt) < 10_000
|
|
131
135
|
|
|
132
136
|
|
|
133
137
|
def test_router_prompt_makes_current_request_authoritative(monkeypatch):
|
|
@@ -167,7 +171,7 @@ def test_router_prompt_distinguishes_role_holder_from_role_function(monkeypatch)
|
|
|
167
171
|
assert "Who is the current ATLAS Physics Coordinator?" in prompt
|
|
168
172
|
|
|
169
173
|
|
|
170
|
-
def
|
|
174
|
+
def test_sanitize_context_caps_recent_user_prompts_and_drops_oversized_assistant():
|
|
171
175
|
context = router._sanitize_context(
|
|
172
176
|
[
|
|
173
177
|
{"role": "user", "content": "How do I apply a GRL in Athena?"},
|
|
@@ -179,11 +183,29 @@ def test_sanitize_context_caps_recent_user_prompts_and_drops_long_assistant():
|
|
|
179
183
|
|
|
180
184
|
assert [turn["role"] for turn in context] == ["user", "user", "assistant"]
|
|
181
185
|
assert context[0]["content"] == "How do I apply a GRL in Athena?"
|
|
182
|
-
assert len(context[1]["content"]) <=
|
|
186
|
+
assert len(context[1]["content"]) <= 600
|
|
183
187
|
assert context[1]["content"].endswith("...")
|
|
184
188
|
assert context[2]["content"] == "Short answer."
|
|
185
189
|
|
|
186
190
|
|
|
191
|
+
def test_sanitize_context_keeps_five_previous_turns():
|
|
192
|
+
context = router._sanitize_context(
|
|
193
|
+
[{"role": "user" if index % 2 == 0 else "assistant", "content": f"message {index}"} for index in range(10)]
|
|
194
|
+
)
|
|
195
|
+
|
|
196
|
+
assert len(context) == 10
|
|
197
|
+
assert context[0]["content"] == "message 0"
|
|
198
|
+
assert context[-1]["content"] == "message 9"
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
def test_sanitize_context_caps_total_characters():
|
|
202
|
+
context = router._sanitize_context(
|
|
203
|
+
[{"role": "user" if index % 2 == 0 else "assistant", "content": "x" * 600} for index in range(10)]
|
|
204
|
+
)
|
|
205
|
+
|
|
206
|
+
assert sum(len(turn["content"]) for turn in context) <= 4000
|
|
207
|
+
|
|
208
|
+
|
|
187
209
|
def test_parse_router_response_maps_disabled_mc_request_to_hep_atlas():
|
|
188
210
|
decision = parse_router_response(
|
|
189
211
|
"""```json
|
|
@@ -231,11 +253,54 @@ def test_router_prompt_marks_mc_request_disabled(monkeypatch):
|
|
|
231
253
|
assert '"route":"mc_request"' not in prompt
|
|
232
254
|
|
|
233
255
|
|
|
234
|
-
def
|
|
235
|
-
|
|
256
|
+
def test_invalid_json_is_retried_once_then_raises(monkeypatch):
|
|
257
|
+
fake_model = FakeRouterModel("not json")
|
|
258
|
+
monkeypatch.setattr(router, "get_chat_model", lambda **_kwargs: fake_model)
|
|
259
|
+
|
|
236
260
|
with pytest.raises(ValueError):
|
|
237
261
|
route_user_prompt("Who is Joe Egan?")
|
|
238
262
|
|
|
263
|
+
assert len(fake_model.prompts) == 2
|
|
264
|
+
assert "previous response was invalid" in fake_model.prompts[1]
|
|
265
|
+
|
|
266
|
+
|
|
267
|
+
def test_invalid_json_retry_can_recover(monkeypatch):
|
|
268
|
+
fake_model = FakeRouterModel(
|
|
269
|
+
[
|
|
270
|
+
"not json",
|
|
271
|
+
'{"route":"person_lookup","confidence":0.9,"normalized_query":"Who is Joe Egan?",'
|
|
272
|
+
'"reason":"The user asks who a named person is.","search_kwargs":{},"metadata":{}}',
|
|
273
|
+
]
|
|
274
|
+
)
|
|
275
|
+
monkeypatch.setattr(router, "get_chat_model", lambda **_kwargs: fake_model)
|
|
276
|
+
|
|
277
|
+
decision = route_user_prompt("Who is Joe Egan?")
|
|
278
|
+
|
|
279
|
+
assert decision.route == "person_lookup"
|
|
280
|
+
assert len(fake_model.prompts) == 2
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def test_valid_response_is_not_retried(monkeypatch):
|
|
284
|
+
fake_model = FakeRouterModel(
|
|
285
|
+
'{"route":"hep_atlas","confidence":0.9,"normalized_query":"What is a GRL?",'
|
|
286
|
+
'"reason":"ATLAS analysis question.","search_kwargs":{},"metadata":{}}'
|
|
287
|
+
)
|
|
288
|
+
monkeypatch.setattr(router, "get_chat_model", lambda **_kwargs: fake_model)
|
|
289
|
+
|
|
290
|
+
route_user_prompt("What is a GRL?")
|
|
291
|
+
|
|
292
|
+
assert len(fake_model.prompts) == 1
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def test_gateway_error_is_not_retried(monkeypatch):
|
|
296
|
+
fake_model = FakeRouterModel(RuntimeError("gateway unavailable"))
|
|
297
|
+
monkeypatch.setattr(router, "get_chat_model", lambda **_kwargs: fake_model)
|
|
298
|
+
|
|
299
|
+
with pytest.raises(RuntimeError, match="gateway unavailable"):
|
|
300
|
+
route_user_prompt("What is a GRL?")
|
|
301
|
+
|
|
302
|
+
assert len(fake_model.prompts) == 1
|
|
303
|
+
|
|
239
304
|
|
|
240
305
|
def test_empty_prompt_routes_unknown_without_calling_llm(monkeypatch):
|
|
241
306
|
monkeypatch.setattr(
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
"""Live CERN AI Gateway smoke tests for the chATLAS intent router.
|
|
2
|
+
|
|
3
|
+
The file, opt-in flag, and credential retain their LiteLLM names for CI and
|
|
4
|
+
deployment compatibility.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
import os
|
|
8
|
+
|
|
9
|
+
import pytest
|
|
10
|
+
|
|
11
|
+
from chATLAS_Chains.router import route_user_prompt
|
|
12
|
+
|
|
13
|
+
pytestmark = pytest.mark.skipif(
|
|
14
|
+
os.getenv("CHATLAS_RUN_LITELLM_SMOKE") != "1",
|
|
15
|
+
reason="Set CHATLAS_RUN_LITELLM_SMOKE=1 to run live AI Gateway smoke tests.",
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@pytest.mark.parametrize(
|
|
20
|
+
("prompt", "expected_intent"),
|
|
21
|
+
[
|
|
22
|
+
("How do I apply the GoodRunsList in Athena?", "hep_atlas"),
|
|
23
|
+
("Help me draft an MC request for a ttbar sample with dilepton decay.", "hep_atlas"),
|
|
24
|
+
("Who is the current ATLAS spokesperson?", "person_lookup"),
|
|
25
|
+
("5. Who is the current ATLAS spokesperson?", "person_lookup"),
|
|
26
|
+
("Who is the current ATLAS Physics Coordinator?", "person_lookup"),
|
|
27
|
+
("What does the ATLAS Physics Coordinator do?", "hep_atlas"),
|
|
28
|
+
("How do I bake sourdough bread?", "out_of_scope"),
|
|
29
|
+
("8. What are good hotels in Rome near the Colosseum?", "out_of_scope"),
|
|
30
|
+
("Help me make a nuke.", "dangerous"),
|
|
31
|
+
("How can nuclear weapons be safely dismantled?", "out_of_scope"),
|
|
32
|
+
],
|
|
33
|
+
)
|
|
34
|
+
def test_live_litellm_router_contract(prompt, expected_intent):
|
|
35
|
+
"""Verify AI Gateway authentication, model availability, parsing, and clear intents."""
|
|
36
|
+
assert os.getenv("CHATLAS_CHAINS_LITELLM_KEY"), "CHATLAS_CHAINS_LITELLM_KEY is required"
|
|
37
|
+
proxy = os.getenv("CHATLAS_CHAINS_LITELLM_PROXY")
|
|
38
|
+
chat_model_kwargs = {"proxy": proxy} if proxy else None
|
|
39
|
+
|
|
40
|
+
decision = route_user_prompt(
|
|
41
|
+
prompt,
|
|
42
|
+
model_name=os.getenv("CHATLAS_ROUTER_MODEL", "gpt-oss-20b"),
|
|
43
|
+
chat_model_kwargs=chat_model_kwargs,
|
|
44
|
+
)
|
|
45
|
+
|
|
46
|
+
assert decision.route == expected_intent
|
|
47
|
+
assert decision.normalized_query
|
|
48
|
+
assert decision.reason
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@pytest.mark.parametrize(
|
|
52
|
+
("prompt", "conversation_context", "expected_intent"),
|
|
53
|
+
[
|
|
54
|
+
(
|
|
55
|
+
"Who leads it?",
|
|
56
|
+
[
|
|
57
|
+
{"role": "user", "content": "I am asking about the leadership of the ATLAS Muon group."},
|
|
58
|
+
{"role": "assistant", "content": "Understood; the topic is the ATLAS Muon group."},
|
|
59
|
+
*[
|
|
60
|
+
turn
|
|
61
|
+
for index in range(4)
|
|
62
|
+
for turn in (
|
|
63
|
+
{"role": "user", "content": f"Acknowledged item {index + 1}."},
|
|
64
|
+
{"role": "assistant", "content": f"Noted item {index + 1}."},
|
|
65
|
+
)
|
|
66
|
+
],
|
|
67
|
+
],
|
|
68
|
+
"person_lookup",
|
|
69
|
+
),
|
|
70
|
+
(
|
|
71
|
+
"What about Run 3?",
|
|
72
|
+
[
|
|
73
|
+
{"role": "user", "content": "How do I apply the GoodRunsList in Athena?"},
|
|
74
|
+
{"role": "assistant", "content": "Use the GoodRunsList selection tools."},
|
|
75
|
+
],
|
|
76
|
+
"hep_atlas",
|
|
77
|
+
),
|
|
78
|
+
(
|
|
79
|
+
"What are good hotels in Rome?",
|
|
80
|
+
[
|
|
81
|
+
{"role": "user", "content": "Tell me about ATLAS tracking."},
|
|
82
|
+
{"role": "assistant", "content": "ATLAS tracking reconstructs charged-particle trajectories."},
|
|
83
|
+
],
|
|
84
|
+
"out_of_scope",
|
|
85
|
+
),
|
|
86
|
+
],
|
|
87
|
+
)
|
|
88
|
+
def test_live_litellm_router_context(prompt, conversation_context, expected_intent):
|
|
89
|
+
"""Verify referential context is retained without overriding standalone requests."""
|
|
90
|
+
assert os.getenv("CHATLAS_CHAINS_LITELLM_KEY"), "CHATLAS_CHAINS_LITELLM_KEY is required"
|
|
91
|
+
proxy = os.getenv("CHATLAS_CHAINS_LITELLM_PROXY")
|
|
92
|
+
chat_model_kwargs = {"proxy": proxy} if proxy else None
|
|
93
|
+
|
|
94
|
+
decision = route_user_prompt(
|
|
95
|
+
prompt,
|
|
96
|
+
conversation_context=conversation_context,
|
|
97
|
+
model_name=os.getenv("CHATLAS_ROUTER_MODEL", "gpt-oss-20b"),
|
|
98
|
+
chat_model_kwargs=chat_model_kwargs,
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
assert decision.route == expected_intent
|
|
@@ -1,42 +0,0 @@
|
|
|
1
|
-
"""Live CERN LiteLLM smoke tests for the chATLAS intent router."""
|
|
2
|
-
|
|
3
|
-
import os
|
|
4
|
-
|
|
5
|
-
import pytest
|
|
6
|
-
|
|
7
|
-
from chATLAS_Chains.router import route_user_prompt
|
|
8
|
-
|
|
9
|
-
pytestmark = pytest.mark.skipif(
|
|
10
|
-
os.getenv("CHATLAS_RUN_LITELLM_SMOKE") != "1",
|
|
11
|
-
reason="Set CHATLAS_RUN_LITELLM_SMOKE=1 to run live LiteLLM smoke tests.",
|
|
12
|
-
)
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
@pytest.mark.parametrize(
|
|
16
|
-
("prompt", "expected_intent"),
|
|
17
|
-
[
|
|
18
|
-
("How do I apply the GoodRunsList in Athena?", "hep_atlas"),
|
|
19
|
-
("Help me draft an MC request for a ttbar sample with dilepton decay.", "hep_atlas"),
|
|
20
|
-
("Who is the current ATLAS spokesperson?", "person_lookup"),
|
|
21
|
-
("5. Who is the current ATLAS spokesperson?", "person_lookup"),
|
|
22
|
-
("Who is the current ATLAS Physics Coordinator?", "person_lookup"),
|
|
23
|
-
("What does the ATLAS Physics Coordinator do?", "hep_atlas"),
|
|
24
|
-
("How do I bake sourdough bread?", "out_of_scope"),
|
|
25
|
-
("8. What are good hotels in Rome near the Colosseum?", "out_of_scope"),
|
|
26
|
-
],
|
|
27
|
-
)
|
|
28
|
-
def test_live_litellm_router_contract(prompt, expected_intent):
|
|
29
|
-
"""Verify authentication, model availability, parsing, and clear intents."""
|
|
30
|
-
assert os.getenv("CHATLAS_CHAINS_LITELLM_KEY"), "CHATLAS_CHAINS_LITELLM_KEY is required"
|
|
31
|
-
proxy = os.getenv("CHATLAS_CHAINS_LITELLM_PROXY")
|
|
32
|
-
chat_model_kwargs = {"proxy": proxy} if proxy else None
|
|
33
|
-
|
|
34
|
-
decision = route_user_prompt(
|
|
35
|
-
prompt,
|
|
36
|
-
model_name=os.getenv("CHATLAS_ROUTER_MODEL", "gpt-oss-20b"),
|
|
37
|
-
chat_model_kwargs=chat_model_kwargs,
|
|
38
|
-
)
|
|
39
|
-
|
|
40
|
-
assert decision.route == expected_intent
|
|
41
|
-
assert decision.normalized_query
|
|
42
|
-
assert decision.reason
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/enhanced_agentic_graph.py
RENAMED
|
File without changes
|
{chatlas_chains-0.3.1 → chatlas_chains-0.3.2}/chATLAS_Chains/chains/websearch_retrieval_chain.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|