adaptive-memory-multi-model-router 2.16.0 → 2.16.2
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.
- package/.github/workflows/adapters-ci.yml +142 -0
- package/.github/workflows/auto-submit-sitemap.yml +41 -0
- package/.github/workflows/ci.yml +2 -5
- package/.github/workflows/mcp-pypi-publish.yml +34 -0
- package/.github/workflows/pypi-publish.yml +146 -0
- package/.github/workflows/tmlpd-publish.yml +23 -0
- package/README.md +245 -148
- package/RELEASE_v2.16.0.md +149 -0
- package/TECHNICAL_README.md +253 -0
- package/adapters/README.md +36 -0
- package/adapters/__init__.py +25 -0
- package/adapters/a3m_adapter/__init__.py +51 -0
- package/adapters/a3m_adapter/adapter/__init__.py +22 -0
- package/adapters/a3m_adapter/adapter/autogen.py +169 -0
- package/adapters/a3m_adapter/adapter/config.py +100 -0
- package/adapters/a3m_adapter/adapter/haystack.py +197 -0
- package/adapters/a3m_adapter/adapter/langchain.py +155 -0
- package/adapters/a3m_adapter/adapter/langgraph.py +196 -0
- package/adapters/a3m_adapter/adapter/llamaindex.py +162 -0
- package/adapters/a3m_adapter/adapter/pinecone.py +217 -0
- package/adapters/a3m_adapter/adapter/vercel.py +188 -0
- package/adapters/a3m_adapter/tests/__init__.py +1 -0
- package/adapters/a3m_adapter/tests/test_adapters.py +118 -0
- package/adapters/a3m_adapter/tests/test_integration.py +80 -0
- package/adapters/requirements-dev.txt +6 -0
- package/adapters/requirements.txt +4 -0
- package/adapters/setup.py +23 -0
- package/demo.py +251 -0
- package/discoverability-diagnosis.md +280 -0
- package/dist/analytics/costAnalytics.d.ts.map +1 -1
- package/dist/benchmark/reproducible.d.ts.map +1 -1
- package/dist/cache/semanticCache.d.ts.map +1 -1
- package/dist/cli/setupWizard.d.ts +257 -50
- package/dist/cli/setupWizard.d.ts.map +1 -1
- package/dist/cli/setupWizard.js +419 -109
- package/dist/cli/setupWizard.js.map +1 -1
- package/dist/cli/tui.d.ts +6 -0
- package/dist/cli/tui.js +96 -67
- package/dist/cli/tui.js.map +1 -0
- package/dist/cli.js +9 -0
- package/dist/cost/budgetEnforcer.d.ts.map +1 -1
- package/dist/cost/costTracker.d.ts.map +1 -1
- package/dist/ensemble/multiRoundDialog.d.ts.map +1 -1
- package/dist/ensemble/shapleyValue.d.ts.map +1 -1
- package/dist/ensemble.d.ts +1 -1
- package/dist/ensemble.js +141 -0
- package/dist/integrations/langchainAdapter.d.ts.map +1 -1
- package/dist/integrations/langchainAdapter.js +3 -3
- package/dist/integrations/langchainAdapter.js.map +1 -1
- package/dist/integrations/oauth.d.ts.map +1 -1
- package/dist/integrations/scienceAdapter.d.ts.map +1 -1
- package/dist/memory/autoFetch.d.ts.map +1 -1
- package/dist/memory/hybridMemory.d.ts.map +1 -1
- package/dist/memory/memoryTree.d.ts.map +1 -1
- package/dist/memory/obsidianVault.d.ts.map +1 -1
- package/dist/memory/reasoningBank.d.ts.map +1 -1
- package/dist/observability/metrics.d.ts.map +1 -1
- package/dist/observability/tracer.d.ts.map +1 -1
- package/dist/providers/providerConfig.d.ts.map +1 -1
- package/dist/providers/providerConfig.js +32 -17
- package/dist/providers/providerConfig.js.map +1 -1
- package/dist/routing/advancedRouter.d.ts.map +1 -1
- package/dist/routing/advancedRouter.js +106 -14
- package/dist/routing/advancedRouter.js.map +1 -1
- package/dist/routing/providerHealth.d.ts.map +1 -1
- package/dist/routing/providerRetry.d.ts.map +1 -1
- package/dist/routing/shadowSampler.d.ts.map +1 -0
- package/dist/routing/shadowSampler.js.map +1 -1
- package/dist/security/guardrails.d.ts.map +1 -1
- package/dist/server/handlers/chatHandler.d.ts.map +1 -1
- package/dist/server/handlers/completionsHandler.d.ts.map +1 -1
- package/dist/server/handlers/embeddingsHandler.d.ts.map +1 -1
- package/dist/server/handlers/healthHandler.d.ts.map +1 -1
- package/dist/server/handlers/metricsHandler.d.ts.map +1 -1
- package/dist/server/handlers/modelsHandler.d.ts.map +1 -1
- package/dist/server/metrics.d.ts.map +1 -1
- package/dist/server/proxyServer.d.ts.map +1 -1
- package/dist/server/router.d.ts.map +1 -1
- package/dist/server/state.d.ts.map +1 -1
- package/dist/skills/__tests__/skill_manager.test.js +5 -265
- package/dist/skills/__tests__/skill_manager.test.js.map +1 -1
- package/dist/utils/tokenUtils.d.ts.map +1 -1
- package/docker-compose.yml +84 -60
- package/docs/ARTICLE_Biology_Inspired_Routing.md +208 -0
- package/docs/ARTICLE_Master.md +78 -0
- package/docs/ARTICLE_Master_CN.md +78 -0
- package/docs/ARTICLE_OpenRouter_Stripe.md +140 -0
- package/docs/DEVPTO_ARTICLE.md +84 -0
- package/docs/HUMAN_STYLE_GUIDE.md +75 -0
- package/docs/IMPRINT_PLAN.md +88 -0
- package/docs/OPENROUTER_ALTERNATIVE.md +184 -0
- package/docs/SOCIAL_CAMPAIGN.md +316 -0
- package/docs/anthropic.html +45 -0
- package/docs/best-llm-routers-2025.html +157 -0
- package/docs/cerebras.html +43 -0
- package/docs/cli-cheatsheet.md +286 -212
- package/docs/comparison.md +2 -2
- package/docs/deepseek.html +44 -0
- package/docs/google.html +47 -0
- package/docs/groq.html +44 -0
- package/docs/llms-full.txt +360 -138
- package/docs/llms.txt +70 -71
- package/docs/mistral.html +43 -0
- package/docs/ollama.html +50 -0
- package/docs/openai.html +57 -0
- package/docs/sitemap.xml +69 -57
- package/docs-site/blog/best-llm-routers-2025.html +157 -0
- package/docs-site/index.html +68 -9
- package/docs-site/providers/anthropic.html +45 -0
- package/docs-site/providers/cerebras.html +43 -0
- package/docs-site/providers/deepseek.html +44 -0
- package/docs-site/providers/google.html +47 -0
- package/docs-site/providers/groq.html +44 -0
- package/docs-site/providers/index.html +41 -0
- package/docs-site/providers/mistral.html +43 -0
- package/docs-site/providers/ollama.html +50 -0
- package/docs-site/providers/openai.html +57 -0
- package/docs-site/sitemap.xml +69 -0
- package/llms.txt +70 -62
- package/package.json +44 -182
- package/packages/agentkit-adapter/LICENSE +21 -0
- package/packages/agentkit-adapter/README.md +126 -0
- package/packages/agentkit-adapter/examples/agentkit-example.ts +139 -0
- package/packages/agentkit-adapter/package.json +57 -0
- package/packages/agentkit-adapter/src/adapter.ts +381 -0
- package/packages/agentkit-adapter/src/index.ts +36 -0
- package/packages/agentkit-adapter/src/types.ts +105 -0
- package/packages/agentkit-adapter/src/util.ts +13 -0
- package/packages/agentkit-adapter/tsconfig.json +22 -0
- package/prometheus.yml +8 -0
- package/python/README.md +35 -81
- package/python/a3m/__init__.py +32 -3
- package/python/a3m/adapters/__init__.py +21 -0
- package/python/a3m/adapters/langchain.py +190 -0
- package/python/a3m/adapters/llamaindex.py +249 -0
- package/python/a3m/adapters/qdrant.py +240 -0
- package/python/a3m/adapters/weaviate.py +263 -0
- package/python/a3m/client.py +5 -0
- package/python/build_verify.sh +32 -0
- package/python/mcp-server/README.md +172 -0
- package/python/mcp-server/a3m_mcp/__init__.py +25 -0
- package/python/mcp-server/a3m_mcp/__main__.py +15 -0
- package/python/mcp-server/a3m_mcp/server.py +205 -0
- package/python/mcp-server/pyproject.toml +24 -0
- package/python/pyproject.toml +60 -6
- package/python/setup.py +3 -28
- package/scripts/submit-sitemap.sh +52 -0
- package/src/__types__/registry.d.ts +14 -0
- package/src/cli/setupWizard.ts +443 -112
- package/src/cli/tui.ts +159 -0
- package/src/ensemble.ts +154 -1
- package/src/integrations/langchainAdapter.ts +2 -2
- package/src/providers/providerConfig.ts +32 -17
- package/src/providers/registry.js +27 -0
- package/src/routing/advancedRouter.ts +99 -14
- package/src/routing/shadowSampler.ts +1 -1
- package/test-install/package.json +12 -0
- package/tests/tsconfig.json +0 -1
- package/tmlpd-pi-extension/README.md +105 -44
- package/tmlpd-pi-extension/docs/demo.svg +33 -0
- package/tmlpd-pi-extension/package.json +35 -106
- package/tmlpd-pi-extension/src/tokenOptimization/contextStratifier.ts +163 -0
- package/tmlpd-pi-extension/src/tokenOptimization/fetchOnceLocal.ts +136 -0
- package/tmlpd-pi-extension/src/tokenOptimization/index.ts +197 -0
- package/tmlpd-pi-extension/src/tokenOptimization/interAgentCompression.ts +157 -0
- package/tmlpd-pi-extension/src/tokenOptimization/schemaContract.ts +101 -0
- package/tmlpd-pi-extension/src/tokenOptimization/semanticCache.ts +248 -0
- package/tmlpd-pi-extension/src/tokenOptimization/tokenAwareFallback.ts +192 -0
- package/tmlpd-pi-extension/test/verify.js +21 -0
- package/tsconfig.build.json +3 -2
- package/packages/a3m-vercel-ai/dist/a3m-language-model.d.ts +0 -12
- package/packages/a3m-vercel-ai/dist/a3m-language-model.d.ts.map +0 -1
- package/packages/a3m-vercel-ai/dist/a3m-language-model.js +0 -289
- package/packages/a3m-vercel-ai/dist/a3m-language-model.js.map +0 -1
- package/packages/a3m-vercel-ai/dist/index.d.ts +0 -82
- package/packages/a3m-vercel-ai/dist/index.d.ts.map +0 -1
- package/packages/a3m-vercel-ai/dist/index.js +0 -79
- package/packages/a3m-vercel-ai/dist/index.js.map +0 -1
- package/packages/a3m-vercel-ai/dist/types.d.ts +0 -97
- package/packages/a3m-vercel-ai/dist/types.d.ts.map +0 -1
- package/packages/a3m-vercel-ai/dist/types.js +0 -5
- package/packages/a3m-vercel-ai/dist/types.js.map +0 -1
- package/src/skills/__tests__/skill_manager.test.ts +0 -328
- package/tmlpd-pi-extension/dist/cache/prefixCache.d.ts +0 -114
- package/tmlpd-pi-extension/dist/cache/prefixCache.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/cache/prefixCache.js +0 -285
- package/tmlpd-pi-extension/dist/cache/prefixCache.js.map +0 -1
- package/tmlpd-pi-extension/dist/cache/responseCache.d.ts +0 -58
- package/tmlpd-pi-extension/dist/cache/responseCache.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/cache/responseCache.js +0 -153
- package/tmlpd-pi-extension/dist/cache/responseCache.js.map +0 -1
- package/tmlpd-pi-extension/dist/cli.js +0 -59
- package/tmlpd-pi-extension/dist/cost/costTracker.d.ts +0 -95
- package/tmlpd-pi-extension/dist/cost/costTracker.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/cost/costTracker.js +0 -240
- package/tmlpd-pi-extension/dist/cost/costTracker.js.map +0 -1
- package/tmlpd-pi-extension/dist/index.d.ts +0 -723
- package/tmlpd-pi-extension/dist/index.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/index.js +0 -239
- package/tmlpd-pi-extension/dist/index.js.map +0 -1
- package/tmlpd-pi-extension/dist/memory/episodicMemory.d.ts +0 -82
- package/tmlpd-pi-extension/dist/memory/episodicMemory.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/memory/episodicMemory.js +0 -145
- package/tmlpd-pi-extension/dist/memory/episodicMemory.js.map +0 -1
- package/tmlpd-pi-extension/dist/orchestration/haloOrchestrator.d.ts +0 -102
- package/tmlpd-pi-extension/dist/orchestration/haloOrchestrator.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/orchestration/haloOrchestrator.js +0 -207
- package/tmlpd-pi-extension/dist/orchestration/haloOrchestrator.js.map +0 -1
- package/tmlpd-pi-extension/dist/orchestration/mctsWorkflow.d.ts +0 -85
- package/tmlpd-pi-extension/dist/orchestration/mctsWorkflow.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/orchestration/mctsWorkflow.js +0 -210
- package/tmlpd-pi-extension/dist/orchestration/mctsWorkflow.js.map +0 -1
- package/tmlpd-pi-extension/dist/providers/localProvider.d.ts +0 -102
- package/tmlpd-pi-extension/dist/providers/localProvider.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/providers/localProvider.js +0 -338
- package/tmlpd-pi-extension/dist/providers/localProvider.js.map +0 -1
- package/tmlpd-pi-extension/dist/providers/registry.d.ts +0 -55
- package/tmlpd-pi-extension/dist/providers/registry.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/providers/registry.js +0 -138
- package/tmlpd-pi-extension/dist/providers/registry.js.map +0 -1
- package/tmlpd-pi-extension/dist/routing/advancedRouter.d.ts +0 -68
- package/tmlpd-pi-extension/dist/routing/advancedRouter.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/routing/advancedRouter.js +0 -332
- package/tmlpd-pi-extension/dist/routing/advancedRouter.js.map +0 -1
- package/tmlpd-pi-extension/dist/tools/tmlpdTools.d.ts +0 -101
- package/tmlpd-pi-extension/dist/tools/tmlpdTools.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/tools/tmlpdTools.js +0 -368
- package/tmlpd-pi-extension/dist/tools/tmlpdTools.js.map +0 -1
- package/tmlpd-pi-extension/dist/utils/batchProcessor.d.ts +0 -96
- package/tmlpd-pi-extension/dist/utils/batchProcessor.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/utils/batchProcessor.js +0 -170
- package/tmlpd-pi-extension/dist/utils/batchProcessor.js.map +0 -1
- package/tmlpd-pi-extension/dist/utils/compression.d.ts +0 -61
- package/tmlpd-pi-extension/dist/utils/compression.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/utils/compression.js +0 -281
- package/tmlpd-pi-extension/dist/utils/compression.js.map +0 -1
- package/tmlpd-pi-extension/dist/utils/reliability.d.ts +0 -74
- package/tmlpd-pi-extension/dist/utils/reliability.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/utils/reliability.js +0 -177
- package/tmlpd-pi-extension/dist/utils/reliability.js.map +0 -1
- package/tmlpd-pi-extension/dist/utils/speculativeDecoding.d.ts +0 -117
- package/tmlpd-pi-extension/dist/utils/speculativeDecoding.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/utils/speculativeDecoding.js +0 -246
- package/tmlpd-pi-extension/dist/utils/speculativeDecoding.js.map +0 -1
- package/tmlpd-pi-extension/dist/utils/tokenUtils.d.ts +0 -50
- package/tmlpd-pi-extension/dist/utils/tokenUtils.d.ts.map +0 -1
- package/tmlpd-pi-extension/dist/utils/tokenUtils.js +0 -124
- package/tmlpd-pi-extension/dist/utils/tokenUtils.js.map +0 -1
|
@@ -0,0 +1,197 @@
|
|
|
1
|
+
"""
|
|
2
|
+
A3M Router Adapter for Haystack (Deepset's RAG framework).
|
|
3
|
+
|
|
4
|
+
Drop-in replacement for Haystack's OpenAIGenerator that routes through A3M Router
|
|
5
|
+
for intelligent, cost-optimized RAG pipelines.
|
|
6
|
+
|
|
7
|
+
Usage:
|
|
8
|
+
from haystack import Pipeline
|
|
9
|
+
from haystack.nodes import Retriever, PromptNode
|
|
10
|
+
from a3m_adapter import A3MHaystackAdapter
|
|
11
|
+
|
|
12
|
+
prompt_node = PromptNode(
|
|
13
|
+
"auto",
|
|
14
|
+
api_key=None,
|
|
15
|
+
generator_type='openai',
|
|
16
|
+
model_adapter=A3MHaystackAdapter(model='auto', parallel_ensemble=2),
|
|
17
|
+
)
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from __future__ import annotations
|
|
21
|
+
|
|
22
|
+
import logging
|
|
23
|
+
from typing import Any, Dict, List, Optional
|
|
24
|
+
|
|
25
|
+
logger = logging.getLogger(__name__)
|
|
26
|
+
|
|
27
|
+
HAYSTACK_AVAILABLE = False
|
|
28
|
+
try:
|
|
29
|
+
from haystack.nodes.base import BaseGenerator
|
|
30
|
+
HAYSTACK_AVAILABLE = True
|
|
31
|
+
except ImportError:
|
|
32
|
+
logger.warning("Haystack not installed. Install with: pip install farm-haystack")
|
|
33
|
+
|
|
34
|
+
A3M_AVAILABLE = False
|
|
35
|
+
try:
|
|
36
|
+
from a3m.router import A3MRouter, RouteResponse
|
|
37
|
+
A3M_AVAILABLE = True
|
|
38
|
+
except ImportError:
|
|
39
|
+
logger.warning(
|
|
40
|
+
"A3M Router not installed. Install with: pip install adaptive-memory-multi-model-router"
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class A3MHaystackAdapter:
|
|
45
|
+
"""
|
|
46
|
+
A3M Router adapter for Haystack's PromptNode.
|
|
47
|
+
|
|
48
|
+
Enables Haystack RAG pipelines to use A3M Router for automatic model selection
|
|
49
|
+
across 47+ providers with cost optimization.
|
|
50
|
+
"""
|
|
51
|
+
|
|
52
|
+
def __init__(
|
|
53
|
+
self,
|
|
54
|
+
model: str = "auto",
|
|
55
|
+
temperature: float = 0.7,
|
|
56
|
+
max_tokens: int = 4096,
|
|
57
|
+
parallel_ensemble: int = 1,
|
|
58
|
+
api_key: Optional[str] = None,
|
|
59
|
+
**kwargs: Any,
|
|
60
|
+
) -> None:
|
|
61
|
+
"""
|
|
62
|
+
Initialize A3M Router adapter for Haystack.
|
|
63
|
+
"""
|
|
64
|
+
self.model = model
|
|
65
|
+
self.temperature = temperature
|
|
66
|
+
self.max_tokens = max_tokens
|
|
67
|
+
self.parallel_ensemble = parallel_ensemble
|
|
68
|
+
self.api_key = api_key
|
|
69
|
+
self._a3m_router = None
|
|
70
|
+
self._initialized = False
|
|
71
|
+
self._kwargs = kwargs
|
|
72
|
+
|
|
73
|
+
def _ensure_router(self) -> None:
|
|
74
|
+
"""Lazily initialize the A3M router."""
|
|
75
|
+
if self._initialized:
|
|
76
|
+
return
|
|
77
|
+
|
|
78
|
+
if not A3M_AVAILABLE:
|
|
79
|
+
raise ImportError(
|
|
80
|
+
"A3M Router is not installed. "
|
|
81
|
+
"Install with: pip install adaptive-memory-multi-model-router"
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
self._a3m_router = A3MRouter(
|
|
85
|
+
model=self.model,
|
|
86
|
+
temperature=self.temperature,
|
|
87
|
+
parallel_ensemble=self.parallel_ensemble,
|
|
88
|
+
)
|
|
89
|
+
self._initialized = True
|
|
90
|
+
logger.info(
|
|
91
|
+
"A3M Router initialized for Haystack: model=%s",
|
|
92
|
+
self.model,
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
def predict(
|
|
96
|
+
self,
|
|
97
|
+
query: str,
|
|
98
|
+
documents: Optional[List[Any]] = None,
|
|
99
|
+
**kwargs: Any,
|
|
100
|
+
) -> Dict[str, Any]:
|
|
101
|
+
"""
|
|
102
|
+
Generate answer from query and optional retrieved documents.
|
|
103
|
+
|
|
104
|
+
Args:
|
|
105
|
+
query: The search query
|
|
106
|
+
documents: Optional list of retrieved documents for RAG
|
|
107
|
+
|
|
108
|
+
Returns:
|
|
109
|
+
Dict with 'answers', 'provider', 'cost'
|
|
110
|
+
"""
|
|
111
|
+
self._ensure_router()
|
|
112
|
+
|
|
113
|
+
# Build context from documents if provided
|
|
114
|
+
if documents:
|
|
115
|
+
context = "\n\n".join([
|
|
116
|
+
f"Document {i+1}: {getattr(doc, 'content', str(doc))}"
|
|
117
|
+
for i, doc in enumerate(documents[:5]) # Limit to 5 docs
|
|
118
|
+
])
|
|
119
|
+
prompt = f"Context:\n{context}\n\nQuestion: {query}\n\nAnswer:"
|
|
120
|
+
else:
|
|
121
|
+
prompt = query
|
|
122
|
+
|
|
123
|
+
messages = [{"role": "user", "content": prompt}]
|
|
124
|
+
|
|
125
|
+
import asyncio
|
|
126
|
+
loop = asyncio.get_event_loop()
|
|
127
|
+
route_result = loop.run_in_executor(
|
|
128
|
+
None,
|
|
129
|
+
lambda: self._a3m_router.route(
|
|
130
|
+
messages=messages,
|
|
131
|
+
temperature=kwargs.get("temperature", self.temperature),
|
|
132
|
+
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
|
133
|
+
**kwargs,
|
|
134
|
+
),
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
return {
|
|
138
|
+
"answers": [{"answer": route_result.content, "score": 1.0}],
|
|
139
|
+
"provider": getattr(route_result, 'provider', 'a3m'),
|
|
140
|
+
"cost": getattr(route_result, 'cost', 0.0),
|
|
141
|
+
}
|
|
142
|
+
|
|
143
|
+
async def apredict(
|
|
144
|
+
self,
|
|
145
|
+
query: str,
|
|
146
|
+
documents: Optional[List[Any]] = None,
|
|
147
|
+
**kwargs: Any,
|
|
148
|
+
) -> Dict[str, Any]:
|
|
149
|
+
"""Async predict for Haystack."""
|
|
150
|
+
self._ensure_router()
|
|
151
|
+
|
|
152
|
+
if documents:
|
|
153
|
+
context = "\n\n".join([
|
|
154
|
+
f"Document {i+1}: {getattr(doc, 'content', str(doc))}"
|
|
155
|
+
for i, doc in enumerate(documents[:5])
|
|
156
|
+
])
|
|
157
|
+
prompt = f"Context:\n{context}\n\nQuestion: {query}\n\nAnswer:"
|
|
158
|
+
else:
|
|
159
|
+
prompt = query
|
|
160
|
+
|
|
161
|
+
messages = [{"role": "user", "content": prompt}]
|
|
162
|
+
|
|
163
|
+
route_result = await self._a3m_router.aroute(
|
|
164
|
+
messages=messages,
|
|
165
|
+
temperature=kwargs.get("temperature", self.temperature),
|
|
166
|
+
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
|
167
|
+
**kwargs,
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
return {
|
|
171
|
+
"answers": [{"answer": route_result.content, "score": 1.0}],
|
|
172
|
+
"provider": getattr(route_result, 'provider', 'a3m'),
|
|
173
|
+
"cost": getattr(route_result, 'cost', 0.0),
|
|
174
|
+
}
|
|
175
|
+
|
|
176
|
+
def run(
|
|
177
|
+
self,
|
|
178
|
+
query: str,
|
|
179
|
+
documents: Optional[List[Any]] = None,
|
|
180
|
+
**kwargs: Any,
|
|
181
|
+
) -> tuple[Dict[str, Any], str]:
|
|
182
|
+
"""
|
|
183
|
+
Haystack-compatible run method.
|
|
184
|
+
|
|
185
|
+
Returns:
|
|
186
|
+
Tuple of (results dict, pipeline run metadata)
|
|
187
|
+
"""
|
|
188
|
+
result = self.predict(query, documents, **kwargs)
|
|
189
|
+
return (result, "a3m-haystack")
|
|
190
|
+
|
|
191
|
+
def __repr__(self) -> str:
|
|
192
|
+
return (
|
|
193
|
+
f"A3MHaystackAdapter("
|
|
194
|
+
f"model={self.model!r}, "
|
|
195
|
+
f"temperature={self.temperature}, "
|
|
196
|
+
f"max_tokens={self.max_tokens})"
|
|
197
|
+
)
|
|
@@ -0,0 +1,155 @@
|
|
|
1
|
+
"""
|
|
2
|
+
A3M Router Adapter for LangChain.
|
|
3
|
+
|
|
4
|
+
Drop-in replacement for LangChain's ChatOpenAI that routes through A3M Router
|
|
5
|
+
for intelligent, cost-optimized model selection across 47+ providers.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
from typing import Any, Dict, List, Optional
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger(__name__)
|
|
14
|
+
|
|
15
|
+
# Check availability
|
|
16
|
+
LANGCHAIN_AVAILABLE = False
|
|
17
|
+
try:
|
|
18
|
+
from langchain_core.language_models import BaseChatModel
|
|
19
|
+
from langchain_core.messages import AIMessage, BaseMessage, HumanMessage, SystemMessage, ToolMessage
|
|
20
|
+
from langchain_core.outputs import ChatGeneration, ChatResult, LLMResult
|
|
21
|
+
LANGCHAIN_AVAILABLE = True
|
|
22
|
+
except ImportError:
|
|
23
|
+
logger.warning("LangChain not installed. Install with: pip install langchain langchain-core")
|
|
24
|
+
|
|
25
|
+
A3M_AVAILABLE = False
|
|
26
|
+
try:
|
|
27
|
+
from a3m.router import A3MRouter, RouteResponse
|
|
28
|
+
A3M_AVAILABLE = True
|
|
29
|
+
except ImportError:
|
|
30
|
+
logger.warning("A3M Router not installed. Install with: pip install adaptive-memory-multi-model-router")
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class A3MLangChainAdapter:
|
|
34
|
+
"""
|
|
35
|
+
A3M Router adapter for LangChain's ChatOpenAI interface.
|
|
36
|
+
|
|
37
|
+
Routes prompts through A3M Router to automatically select the cheapest
|
|
38
|
+
capable model across 47+ LLM providers.
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(
|
|
42
|
+
self,
|
|
43
|
+
model: str = "auto",
|
|
44
|
+
temperature: float = 0.0,
|
|
45
|
+
max_tokens: Optional[int] = 4096,
|
|
46
|
+
parallel_ensemble: int = 1,
|
|
47
|
+
api_key: Optional[str] = None,
|
|
48
|
+
**kwargs: Any,
|
|
49
|
+
) -> None:
|
|
50
|
+
"""
|
|
51
|
+
Initialize A3M Router adapter.
|
|
52
|
+
|
|
53
|
+
Args:
|
|
54
|
+
model: Model name or "auto" for automatic routing
|
|
55
|
+
temperature: Sampling temperature
|
|
56
|
+
max_tokens: Maximum tokens to generate
|
|
57
|
+
parallel_ensemble: Number of providers to run in parallel
|
|
58
|
+
api_key: A3M API key (optional)
|
|
59
|
+
"""
|
|
60
|
+
self.model = model
|
|
61
|
+
self.temperature = temperature
|
|
62
|
+
self.max_tokens = max_tokens
|
|
63
|
+
self.parallel_ensemble = parallel_ensemble
|
|
64
|
+
self.api_key = api_key
|
|
65
|
+
self._a3m_router = None
|
|
66
|
+
self._initialized = False
|
|
67
|
+
|
|
68
|
+
def _ensure_router(self) -> None:
|
|
69
|
+
"""Lazily initialize the A3M router."""
|
|
70
|
+
if self._initialized:
|
|
71
|
+
return
|
|
72
|
+
|
|
73
|
+
if not A3M_AVAILABLE:
|
|
74
|
+
raise ImportError(
|
|
75
|
+
"A3M Router is not installed. "
|
|
76
|
+
"Install with: pip install adaptive-memory-multi-model-router"
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
self._a3m_router = A3MRouter(
|
|
80
|
+
model=self.model,
|
|
81
|
+
temperature=self.temperature,
|
|
82
|
+
parallel_ensemble=self.parallel_ensemble,
|
|
83
|
+
)
|
|
84
|
+
self._initialized = True
|
|
85
|
+
logger.info(
|
|
86
|
+
"A3M Router initialized: model=%s, ensemble=%d",
|
|
87
|
+
self.model,
|
|
88
|
+
self.parallel_ensemble,
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
@property
|
|
92
|
+
def _llm_type(self) -> str:
|
|
93
|
+
return "a3m_router"
|
|
94
|
+
|
|
95
|
+
def _generate(
|
|
96
|
+
self,
|
|
97
|
+
messages: List[BaseMessage],
|
|
98
|
+
stop: Optional[List[str]] = None,
|
|
99
|
+
run_manager: Any = None,
|
|
100
|
+
**kwargs: Any,
|
|
101
|
+
) -> LLMResult:
|
|
102
|
+
"""Generate a response using A3M Router."""
|
|
103
|
+
self._ensure_router()
|
|
104
|
+
|
|
105
|
+
# Convert messages
|
|
106
|
+
a3m_messages = self._convert_messages(messages)
|
|
107
|
+
|
|
108
|
+
# Route through A3M
|
|
109
|
+
import asyncio
|
|
110
|
+
loop = asyncio.get_event_loop()
|
|
111
|
+
route_result = loop.run_in_executor(
|
|
112
|
+
None,
|
|
113
|
+
lambda: self._a3m_router.route(
|
|
114
|
+
messages=a3m_messages,
|
|
115
|
+
temperature=self.temperature,
|
|
116
|
+
max_tokens=self.max_tokens,
|
|
117
|
+
stop=stop,
|
|
118
|
+
**kwargs,
|
|
119
|
+
),
|
|
120
|
+
)
|
|
121
|
+
|
|
122
|
+
ai_message = AIMessage(content=route_result.content)
|
|
123
|
+
generation = ChatGeneration(message=ai_message)
|
|
124
|
+
return LLMResult(generations=[[generation]])
|
|
125
|
+
|
|
126
|
+
def _convert_messages(self, messages: List[BaseMessage]) -> List[Dict[str, Any]]:
|
|
127
|
+
"""Convert LangChain messages to A3M format."""
|
|
128
|
+
a3m_messages = []
|
|
129
|
+
for msg in messages:
|
|
130
|
+
if isinstance(msg, SystemMessage):
|
|
131
|
+
a3m_messages.append({"role": "system", "content": msg.content})
|
|
132
|
+
elif isinstance(msg, HumanMessage):
|
|
133
|
+
a3m_messages.append({"role": "user", "content": msg.content})
|
|
134
|
+
elif isinstance(msg, AIMessage):
|
|
135
|
+
a3m_messages.append({"role": "assistant", "content": msg.content})
|
|
136
|
+
elif isinstance(msg, ToolMessage):
|
|
137
|
+
a3m_messages.append(
|
|
138
|
+
{"role": "tool", "content": msg.content, "tool_call_id": msg.tool_call_id}
|
|
139
|
+
)
|
|
140
|
+
else:
|
|
141
|
+
a3m_messages.append({"role": "user", "content": str(msg)})
|
|
142
|
+
return a3m_messages
|
|
143
|
+
|
|
144
|
+
def bind_tools(self, tools: List[Dict[str, Any]], **kwargs: Any) -> "A3MLangChainAdapter":
|
|
145
|
+
"""Bind tools for function calling."""
|
|
146
|
+
return self
|
|
147
|
+
|
|
148
|
+
def __repr__(self) -> str:
|
|
149
|
+
return (
|
|
150
|
+
f"A3MLangChainAdapter("
|
|
151
|
+
f"model={self.model!r}, "
|
|
152
|
+
f"temperature={self.temperature}, "
|
|
153
|
+
f"max_tokens={self.max_tokens}, "
|
|
154
|
+
f"ensemble={self.parallel_ensemble})"
|
|
155
|
+
)
|
|
@@ -0,0 +1,196 @@
|
|
|
1
|
+
"""
|
|
2
|
+
A3M Router Adapter for LangGraph (Microsoft's agent framework).
|
|
3
|
+
|
|
4
|
+
Drop-in replacement for LangGraph's stateful agent that routes through A3M Router
|
|
5
|
+
for intelligent, cost-optimized multi-step conversations.
|
|
6
|
+
|
|
7
|
+
Usage:
|
|
8
|
+
from langgraph.prebuilt import create_react_agent
|
|
9
|
+
from a3m_adapter import A3MLangGraphAdapter
|
|
10
|
+
|
|
11
|
+
adapter = A3MLangGraphAdapter(model='auto', parallel_ensemble=2)
|
|
12
|
+
|
|
13
|
+
agent = create_react_agent(adapter, tools=[...])
|
|
14
|
+
|
|
15
|
+
result = agent.invoke({"messages": [{"role": "user", "content": "Hello"]})
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import logging
|
|
21
|
+
from typing import Any, Dict, List, TypedDict
|
|
22
|
+
|
|
23
|
+
logger = logging.getLogger(__name__)
|
|
24
|
+
|
|
25
|
+
LANGGRAPH_AVAILABLE = False
|
|
26
|
+
try:
|
|
27
|
+
import langgraph
|
|
28
|
+
from langgraph.prebuilt import create_react_agent
|
|
29
|
+
from langchain_core.messages import BaseMessage, AIMessage, HumanMessage
|
|
30
|
+
LANGGRAPH_AVAILABLE = True
|
|
31
|
+
except ImportError:
|
|
32
|
+
logger.warning("LangGraph not installed. Install with: pip install langgraph")
|
|
33
|
+
|
|
34
|
+
A3M_AVAILABLE = False
|
|
35
|
+
try:
|
|
36
|
+
from a3m.router import A3MRouter, RouteResponse
|
|
37
|
+
A3M_AVAILABLE = True
|
|
38
|
+
except ImportError:
|
|
39
|
+
logger.warning(
|
|
40
|
+
"A3M Router not installed. Install with: pip install adaptive-memory-multi-model-router"
|
|
41
|
+
)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class A3MLangGraphAdapter:
|
|
45
|
+
"""
|
|
46
|
+
A3M Router adapter for LangGraph's prebuilt agents.
|
|
47
|
+
|
|
48
|
+
Enables LangGraph agents to use A3M Router for automatic model selection
|
|
49
|
+
across 47+ providers with cost optimization and stateful conversations.
|
|
50
|
+
"""
|
|
51
|
+
|
|
52
|
+
def __init__(
|
|
53
|
+
self,
|
|
54
|
+
model: str = "auto",
|
|
55
|
+
temperature: float = 0.7,
|
|
56
|
+
max_tokens: int = 4096,
|
|
57
|
+
parallel_ensemble: int = 1,
|
|
58
|
+
api_key: Optional[str] = None,
|
|
59
|
+
**kwargs: Any,
|
|
60
|
+
) -> None:
|
|
61
|
+
"""
|
|
62
|
+
Initialize A3M Router adapter for LangGraph.
|
|
63
|
+
"""
|
|
64
|
+
self.model = model
|
|
65
|
+
self.temperature = temperature
|
|
66
|
+
self.max_tokens = max_tokens
|
|
67
|
+
self.parallel_ensemble = parallel_ensemble
|
|
68
|
+
self.api_key = api_key
|
|
69
|
+
self._a3m_router = None
|
|
70
|
+
self._initialized = False
|
|
71
|
+
self._kwargs = kwargs
|
|
72
|
+
|
|
73
|
+
def _ensure_router(self) -> None:
|
|
74
|
+
"""Lazily initialize the A3M router."""
|
|
75
|
+
if self._initialized:
|
|
76
|
+
return
|
|
77
|
+
|
|
78
|
+
if not A3M_AVAILABLE:
|
|
79
|
+
raise ImportError(
|
|
80
|
+
"A3M Router is not installed. "
|
|
81
|
+
"Install with: pip install adaptive-memory-multi-model-router"
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
self._a3m_router = A3MRouter(
|
|
85
|
+
model=self.model,
|
|
86
|
+
temperature=self.temperature,
|
|
87
|
+
parallel_ensemble=self.parallel_ensemble,
|
|
88
|
+
)
|
|
89
|
+
self._initialized = True
|
|
90
|
+
logger.info(
|
|
91
|
+
"A3M Router initialized for LangGraph: model=%s, ensemble=%d",
|
|
92
|
+
self.model,
|
|
93
|
+
self.parallel_ensemble,
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
def get_model(self):
|
|
97
|
+
"""
|
|
98
|
+
Get the underlying model for LangGraph.
|
|
99
|
+
|
|
100
|
+
Returns an object compatible with LangGraph's prebuilt agents.
|
|
101
|
+
"""
|
|
102
|
+
self._ensure_router()
|
|
103
|
+
return self
|
|
104
|
+
|
|
105
|
+
def __call__(
|
|
106
|
+
self,
|
|
107
|
+
state: Dict[str, Any],
|
|
108
|
+
**kwargs: Any,
|
|
109
|
+
) -> Dict[str, Any]:
|
|
110
|
+
"""
|
|
111
|
+
LangGraph-compatible callable for node execution.
|
|
112
|
+
|
|
113
|
+
Args:
|
|
114
|
+
state: LangGraph state dict with 'messages' key
|
|
115
|
+
|
|
116
|
+
Returns:
|
|
117
|
+
Updated state dict
|
|
118
|
+
"""
|
|
119
|
+
self._ensure_router()
|
|
120
|
+
|
|
121
|
+
messages = state.get("messages", [])
|
|
122
|
+
|
|
123
|
+
# Convert LangGraph messages to A3M format
|
|
124
|
+
a3m_messages = self._convert_messages(messages)
|
|
125
|
+
|
|
126
|
+
import asyncio
|
|
127
|
+
loop = asyncio.get_event_loop()
|
|
128
|
+
route_result = loop.run_in_executor(
|
|
129
|
+
None,
|
|
130
|
+
lambda: self._a3m_router.route(
|
|
131
|
+
messages=a3m_messages,
|
|
132
|
+
temperature=kwargs.get("temperature", self.temperature),
|
|
133
|
+
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
|
134
|
+
**kwargs,
|
|
135
|
+
),
|
|
136
|
+
)
|
|
137
|
+
|
|
138
|
+
# Add response to messages
|
|
139
|
+
new_messages = messages + [
|
|
140
|
+
AIMessage(content=route_result.content)
|
|
141
|
+
]
|
|
142
|
+
|
|
143
|
+
return {
|
|
144
|
+
**state,
|
|
145
|
+
"messages": new_messages,
|
|
146
|
+
}
|
|
147
|
+
|
|
148
|
+
def _convert_messages(
|
|
149
|
+
self,
|
|
150
|
+
messages: List[BaseMessage],
|
|
151
|
+
) -> List[Dict[str, Any]]:
|
|
152
|
+
"""Convert LangGraph messages to A3M format."""
|
|
153
|
+
a3m_messages = []
|
|
154
|
+
for msg in messages:
|
|
155
|
+
if isinstance(msg, HumanMessage):
|
|
156
|
+
a3m_messages.append({"role": "user", "content": msg.content})
|
|
157
|
+
elif isinstance(msg, AIMessage):
|
|
158
|
+
a3m_messages.append({"role": "assistant", "content": msg.content})
|
|
159
|
+
else:
|
|
160
|
+
a3m_messages.append({"role": "user", "content": str(msg)})
|
|
161
|
+
return a3m_messages
|
|
162
|
+
|
|
163
|
+
async def ainvoke(
|
|
164
|
+
self,
|
|
165
|
+
state: Dict[str, Any],
|
|
166
|
+
**kwargs: Any,
|
|
167
|
+
) -> Dict[str, Any]:
|
|
168
|
+
"""Async version of __call__."""
|
|
169
|
+
self._ensure_router()
|
|
170
|
+
|
|
171
|
+
messages = state.get("messages", [])
|
|
172
|
+
a3m_messages = self._convert_messages(messages)
|
|
173
|
+
|
|
174
|
+
route_result = await self._a3m_router.aroute(
|
|
175
|
+
messages=a3m_messages,
|
|
176
|
+
temperature=kwargs.get("temperature", self.temperature),
|
|
177
|
+
max_tokens=kwargs.get("max_tokens", self.max_tokens),
|
|
178
|
+
**kwargs,
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
new_messages = messages + [
|
|
182
|
+
AIMessage(content=route_result.content)
|
|
183
|
+
]
|
|
184
|
+
|
|
185
|
+
return {
|
|
186
|
+
**state,
|
|
187
|
+
"messages": new_messages,
|
|
188
|
+
}
|
|
189
|
+
|
|
190
|
+
def __repr__(self) -> str:
|
|
191
|
+
return (
|
|
192
|
+
f"A3MLangGraphAdapter("
|
|
193
|
+
f"model={self.model!r}, "
|
|
194
|
+
f"temperature={self.temperature}, "
|
|
195
|
+
f"ensemble={self.parallel_ensemble})"
|
|
196
|
+
)
|
|
@@ -0,0 +1,162 @@
|
|
|
1
|
+
"""
|
|
2
|
+
A3M Router Adapter for LlamaIndex.
|
|
3
|
+
|
|
4
|
+
Drop-in replacement for LlamaIndex's BaseLLM that routes through A3M Router
|
|
5
|
+
for intelligent, cost-optimized model selection.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import logging
|
|
11
|
+
from typing import Any, Dict, List, Optional, Sequence
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger(__name__)
|
|
14
|
+
|
|
15
|
+
# Check availability
|
|
16
|
+
LLAMAINDEX_AVAILABLE = False
|
|
17
|
+
_llama_metadata_class = None
|
|
18
|
+
try:
|
|
19
|
+
from llama_index.core.base.llms.base import BaseLLM, CompletionResponse
|
|
20
|
+
from llama_index.core.base.llms.types import ChatMessage
|
|
21
|
+
LLAMAINDEX_AVAILABLE = True
|
|
22
|
+
try:
|
|
23
|
+
from llama_index.core.base.llms.base import LLMMetadata
|
|
24
|
+
_llama_metadata_class = LLMMetadata
|
|
25
|
+
except ImportError:
|
|
26
|
+
pass
|
|
27
|
+
except ImportError:
|
|
28
|
+
logger.warning("LlamaIndex not installed. Install with: pip install llama-index")
|
|
29
|
+
|
|
30
|
+
A3M_AVAILABLE = False
|
|
31
|
+
try:
|
|
32
|
+
from a3m.router import A3MRouter, RouteResponse
|
|
33
|
+
A3M_AVAILABLE = True
|
|
34
|
+
except ImportError:
|
|
35
|
+
logger.warning("A3M Router not installed. Install with: pip install adaptive-memory-multi-model-router")
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class A3MLlamaIndexAdapter:
|
|
39
|
+
"""
|
|
40
|
+
A3M Router adapter for LlamaIndex's BaseLLM interface.
|
|
41
|
+
|
|
42
|
+
Routes prompts through A3M Router to automatically select the cheapest
|
|
43
|
+
capable model across 47+ LLM providers.
|
|
44
|
+
"""
|
|
45
|
+
|
|
46
|
+
def __init__(
|
|
47
|
+
self,
|
|
48
|
+
model: str = "auto",
|
|
49
|
+
temperature: float = 0.0,
|
|
50
|
+
max_tokens: Optional[int] = 4096,
|
|
51
|
+
parallel_ensemble: int = 1,
|
|
52
|
+
api_key: Optional[str] = None,
|
|
53
|
+
**kwargs: Any,
|
|
54
|
+
) -> None:
|
|
55
|
+
"""
|
|
56
|
+
Initialize A3M Router adapter.
|
|
57
|
+
"""
|
|
58
|
+
self.model = model
|
|
59
|
+
self.temperature = temperature
|
|
60
|
+
self.max_tokens = max_tokens
|
|
61
|
+
self.parallel_ensemble = parallel_ensemble
|
|
62
|
+
self.api_key = api_key
|
|
63
|
+
self._a3m_router = None
|
|
64
|
+
self._initialized = False
|
|
65
|
+
|
|
66
|
+
def _ensure_router(self) -> None:
|
|
67
|
+
"""Lazily initialize the A3M router."""
|
|
68
|
+
if self._initialized:
|
|
69
|
+
return
|
|
70
|
+
|
|
71
|
+
if not A3M_AVAILABLE:
|
|
72
|
+
raise ImportError(
|
|
73
|
+
"A3M Router is not installed. "
|
|
74
|
+
"Install with: pip install adaptive-memory-multi-model-router"
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
self._a3m_router = A3MRouter(
|
|
78
|
+
model=self.model,
|
|
79
|
+
temperature=self.temperature,
|
|
80
|
+
parallel_ensemble=self.parallel_ensemble,
|
|
81
|
+
)
|
|
82
|
+
self._initialized = True
|
|
83
|
+
logger.info(
|
|
84
|
+
"A3M Router initialized: model=%s, ensemble=%d",
|
|
85
|
+
self.model,
|
|
86
|
+
self.parallel_ensemble,
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
@property
|
|
90
|
+
def metadata(self) -> Dict[str, Any]:
|
|
91
|
+
"""Return LLM metadata as a dict (framework-agnostic)."""
|
|
92
|
+
return {
|
|
93
|
+
"context_window": 128000,
|
|
94
|
+
"num_output": self.max_tokens or 4096,
|
|
95
|
+
"model_name": self.model,
|
|
96
|
+
"is_chat_model": True,
|
|
97
|
+
}
|
|
98
|
+
|
|
99
|
+
def complete(self, prompt: str, **kwargs: Any) -> CompletionResponse:
|
|
100
|
+
"""Complete a prompt using A3M Router."""
|
|
101
|
+
self._ensure_router()
|
|
102
|
+
|
|
103
|
+
messages = [{"role": "user", "content": prompt}]
|
|
104
|
+
|
|
105
|
+
import asyncio
|
|
106
|
+
loop = asyncio.get_event_loop()
|
|
107
|
+
route_result = loop.run_in_executor(
|
|
108
|
+
None,
|
|
109
|
+
lambda: self._a3m_router.route(
|
|
110
|
+
messages=messages,
|
|
111
|
+
temperature=self.temperature,
|
|
112
|
+
max_tokens=self.max_tokens,
|
|
113
|
+
**kwargs,
|
|
114
|
+
),
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
return CompletionResponse(text=route_result.content, raw=route_result)
|
|
118
|
+
|
|
119
|
+
def chat(self, messages: Sequence[ChatMessage], **kwargs: Any) -> CompletionResponse:
|
|
120
|
+
"""Chat completion using A3M Router."""
|
|
121
|
+
self._ensure_router()
|
|
122
|
+
|
|
123
|
+
a3m_messages = self._convert_messages(messages)
|
|
124
|
+
|
|
125
|
+
import asyncio
|
|
126
|
+
loop = asyncio.get_event_loop()
|
|
127
|
+
route_result = loop.run_in_executor(
|
|
128
|
+
None,
|
|
129
|
+
lambda: self._a3m_router.route(
|
|
130
|
+
messages=a3m_messages,
|
|
131
|
+
temperature=self.temperature,
|
|
132
|
+
max_tokens=self.max_tokens,
|
|
133
|
+
**kwargs,
|
|
134
|
+
),
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
return CompletionResponse(text=route_result.content, raw=route_result)
|
|
138
|
+
|
|
139
|
+
def _convert_messages(self, messages: Sequence[ChatMessage]) -> List[Dict[str, Any]]:
|
|
140
|
+
"""Convert LlamaIndex ChatMessages to A3M format."""
|
|
141
|
+
a3m_messages = []
|
|
142
|
+
for msg in messages:
|
|
143
|
+
role = msg.role.value if hasattr(msg.role, 'value') else str(msg.role).lower()
|
|
144
|
+
role_map = {
|
|
145
|
+
"system": "system",
|
|
146
|
+
"user": "user",
|
|
147
|
+
"assistant": "assistant",
|
|
148
|
+
"tool": "tool",
|
|
149
|
+
"function": "function",
|
|
150
|
+
}
|
|
151
|
+
a3m_role = role_map.get(role, "user")
|
|
152
|
+
a3m_messages.append({"role": a3m_role, "content": msg.content})
|
|
153
|
+
return a3m_messages
|
|
154
|
+
|
|
155
|
+
def __repr__(self) -> str:
|
|
156
|
+
return (
|
|
157
|
+
f"A3MLlamaIndexAdapter("
|
|
158
|
+
f"model={self.model!r}, "
|
|
159
|
+
f"temperature={self.temperature}, "
|
|
160
|
+
f"max_tokens={self.max_tokens}, "
|
|
161
|
+
f"ensemble={self.parallel_ensemble})"
|
|
162
|
+
)
|