Skip to content

Commit 0139f60

Browse files
fix(langchain): send configured authentication headers to LangServe (#202)
Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> Co-authored-by: n-papaioannou <n.papaioannou@ime.life>
1 parent a27b238 commit 0139f60

2 files changed

Lines changed: 34 additions & 1 deletion

File tree

‎ifixai/providers/langchain.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
ProviderResponseError,
1616
ProviderTimeoutError,
1717
)
18+
from ifixai.providers.http import _build_auth_headers
1819

1920
DEFAULT_ENDPOINT = "http://localhost:8000"
2021

@@ -44,7 +45,7 @@ async def send_message(
4445

4546
try:
4647
async with aiohttp.ClientSession(timeout=timeout) as session:
47-
async with session.post(url, json=payload) as response:
48+
async with session.post(url, json=payload, headers=_build_auth_headers(config)) as response:
4849
if response.status == 401 or response.status == 403:
4950
raise ProviderAuthError(
5051
provider="langchain",

‎ifixai/tests/test_langchain_transport_native.py‎

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -103,3 +103,35 @@ async def test_disconnected_endpoint_remains_connection_error():
103103
await LangChainProvider().send_message(
104104
[ChatMessage(role="user", content="owned prompt")], config(url)
105105
)
106+
107+
108+
@pytest.mark.asyncio
109+
@pytest.mark.parametrize("auth,key,custom,required", [
110+
("bearer", "owned-key", {}, {"Authorization": "Bearer owned-key"}),
111+
("basic", "owned:password", {}, {"Authorization": "Basic b3duZWQ6cGFzc3dvcmQ="}),
112+
("api_key", "owned-key", {}, {"X-API-Key": "owned-key"}),
113+
("none", "", {"X-Owned-Token": "owned-token"}, {"X-Owned-Token": "owned-token"}),
114+
])
115+
async def test_langserve_authenticated_invoke_uses_configured_headers(monkeypatch, auth, key, custom, required):
116+
monkeypatch.delenv("IFIXAI_EXTRA_HEADERS", raising=False)
117+
seen = []
118+
async def invoke(request):
119+
seen.append(dict(request.headers))
120+
if any(request.headers.get(name) != value for name, value in required.items()):
121+
return web.json_response({"detail": "owned authentication required"}, status=401)
122+
return web.json_response({"output": "owned authenticated reply"})
123+
124+
app = web.Application()
125+
app.router.add_post("/invoke", invoke)
126+
runner = web.AppRunner(app)
127+
await runner.setup()
128+
site = web.TCPSite(runner, "127.0.0.1", 0)
129+
await site.start()
130+
url = f"http://127.0.0.1:{site._server.sockets[0].getsockname()[1]}"
131+
try:
132+
config = ProviderConfig(provider="langchain", endpoint=url, auth_method=auth,
133+
api_key=key, extra_headers=custom, max_retries=0)
134+
assert await LangChainProvider().send_message([ChatMessage(content="hello")], config) == "owned authenticated reply"
135+
assert len(seen) == 1
136+
finally:
137+
await runner.cleanup()

0 commit comments

Comments
 (0)