Implementing a Custom Chat Model with LangChain

This guide provides a comprehensive blueprint for creating a custom chat model by subclassing LangChain's BaseChatModel, including configuration, method overrides, and error handling.

Blog cover image
2101050's avatar
2101050
77 views

Implementation Blueprint: Required Overrides and Configuration

This section lays out the minimal interface contract you must fulfill to wrap any external chat API under LangChain’s BaseChatModel abstraction.

1.1 Core Class Signature

Your custom model subclass must inherit from langchain_core.language_models.BaseChatModel. For example, in v0.3 of the docs:

Python
1from langchain_core.language_models import BaseChatModel 2from pydantic import Field 3 4class MyCustomChatModel(BaseChatModel): 5 """Wraps Acme’s chat endpoint under LangChain.” 6 7 api_key: str 8 endpoint_url: str 9 model_name: str = Field(alias="model") 10 temperature: float = Field(default=0.7) 11 max_tokens: int = Field(default=512) 12 timeout: Optional[int] = None 13 stop: Optional[List[str]] = None 14 max_retries: int = Field(default=3) 15
  • Pydantic fields become the constructor args (e.g., MyCustomChatModel(model="x", api_key="…")).
  • Use Field(alias="model") if your API expects a different key.

1.2 Mandatory Method Overrides

At minimum, override:

  1. _generate(self, messages, stop=None, run_manager=None, **kwargs) -> ChatResult

  2. _agenerate(self, messages, stop=None, run_manager=None, **kwargs) -> ChatResult

  3. _stream(self, messages, stop=None, run_manager=None, **kwargs) -> Iterator[ChatGenerationChunk] (only if your API supports streaming)

  4. _identifying_params(self) -> Dict[str, Any]

The v0.3 guide’s ChatParrotLink example shows a skeleton:

Python
1from langchain_core.callbacks import CallbackManagerForLLMRun 2from langchain_core.messages import BaseMessage, AIMessage 3from langchain_core.outputs import ChatGeneration, ChatResult 4 5class ChatParrotLink(BaseChatModel): 6 model_name: str = Field(alias="model") 7 parrot_buffer_length: int 8 # ... other config fields ... 9 10 def _generate(self, messages: List[BaseMessage], stop=None, run_manager=None, **kwargs) -> ChatResult: 11 # Step 1: call run_manager.on_llm_start 12 run_manager.on_llm_start([m.content for m in messages], **self._identifying_params()) 13 # Step 2: send HTTP request to your API: 14 resp = requests.post(self.endpoint_url, json={…}, headers={…}) 15 text = resp.json()["choices"][0]["text"] 16 # Step 3: wrap into ChatGeneration/ChatResult 17 generation = ChatGeneration( 18 message=AIMessage(content=text), 19 generation_info={"model": self.model_name}, 20 # optional: usage_metadata={"prompt_tokens": …, "completion_tokens": …} 21 ) 22 run_manager.on_llm_end([generation], **self._identifying_params()) 23 return ChatResult(generations=[generation]) 24 25 async def _agenerate(self, messages, stop=None, run_manager=None, **kwargs) -> ChatResult: 26 # same as _generate but using httpx.AsyncClient and await 27 … 28 29 def _stream(self, messages, stop=None, run_manager=None, **kwargs): 30 # if streaming supported, yield ChatGenerationChunk 31 … 32 33 def _identifying_params(self) -> Dict[str, Any]: 34 return {"model_name": self.model_name, "temperature": self.temperature} 35

1.3 Callback Integration Checklist

  • run_manager.on_llm_start before sending prompt
  • run_manager.on_llm_stream for each token/chunk (in streaming)
  • run_manager.on_llm_end after final response
  • For async: use self.astream_events to emit events

1.4 Packaging Responses

  • Sync: return ChatResult(generations=[ChatGeneration(… )])
  • Async: same return type; LangChain adapts ainvoke
  • Streaming: iterate and yield ChatGenerationChunk objects, containing partial AIMessageChunk

1.5 Identifying Params

Ensure _identifying_params returns a dict of all config fields that define which underlying model variant you’re calling (e.g., model name, temperature). LangChain uses this for hashing and caching.


2. Subclassing BaseChatModel: Sync & Async Methods

Now we build the custom wrapper step by step, focusing on synchronous and asynchronous path implementations.

2.1 Defining Configuration Fields

Python
1from pydantic import Field 2from typing import Optional, List 3 4class AcmeChatModel(BaseChatModel): 5 api_key: str = Field(..., description="Acme Cloud API key") 6 endpoint_url: str = Field("https://api.acme.ai/v1/chat", description="API endpoint") 7 model_name: str = Field(alias="model", description="Model name, e.g., acme-chat-turbo") 8 temperature: float = Field(default=0.7, description="Sampling temperature") 9 max_tokens: int = Field(default=512, description="Output token limit") 10 timeout: Optional[int] = Field(default=30, description="Request timeout in seconds") 11 stop: Optional[List[str]] = Field(default=None, description="Stop sequences") 12 max_retries: int = Field(default=3, description="Number of retry attempts") 13
  • Type hints and Field descriptions guarantee correct initialization and auto-documentation in LangChain’s CLI or notebooks.

2.2 Implementing _generate

Python
1import requests 2from langchain_core.messages import BaseMessage, AIMessage 3from langchain_core.outputs import ChatResult, ChatGeneration 4from langchain_core.callbacks import CallbackManagerForLLMRun 5 6 def _generate( 7 self, 8 messages: List[BaseMessage], 9 stop: Optional[List[str]] = None, 10 run_manager: Optional[CallbackManagerForLLMRun] = None, 11 **kwargs, 12 ) -> ChatResult: 13 # 1. Emit start event 14 inputs = [msg.content for msg in messages] 15 run_manager.on_llm_start(inputs, **self._identifying_params()) 16 17 # 2. Build payload 18 payload = { 19 "model": self.model_name, 20 "messages": [{"role": m.type, "content": m.content} for m in messages], 21 "temperature": self.temperature, 22 "max_tokens": self.max_tokens, 23 **({"stop": stop} if stop else {}), 24 } 25 26 # 3. HTTP call with retries 27 for attempt in range(self.max_retries): 28 try: 29 resp = requests.post( 30 self.endpoint_url, 31 json=payload, 32 headers={"Authorization": f"Bearer {self.api_key}"}, 33 timeout=self.timeout, 34 ) 35 resp.raise_for_status() 36 break 37 except Exception as e: 38 if attempt + 1 == self.max_retries: 39 raise 40 data = resp.json() 41 42 # 4. Extract output 43 text = data["choices"][0]["message"]["content"] 44 usage = data.get("usage", {}) 45 46 # 5. Wrap into ChatGeneration 47 generation = ChatGeneration( 48 message=AIMessage(content=text), 49 generation_info={"model": self.model_name}, 50 usage=usage, 51 ) 52 53 # 6. Emit end event 54 run_manager.on_llm_end([generation], **self._identifying_params()) 55 56 return ChatResult(generations=[generation]) 57

Notes:

  • Collect token counts from usage if available, else compute via string length.
  • Include **self._identifying_params() in callback events.

2.3 Implementing _agenerate

Python
1import httpx 2import asyncio 3 4 async def _agenerate( 5 self, 6 messages: List[BaseMessage], 7 stop: Optional[List[str]] = None, 8 run_manager: Optional[CallbackManagerForLLMRun] = None, 9 **kwargs, 10 ) -> ChatResult: 11 # Async start event 12 inputs = [m.content for m in messages] 13 await run_manager.on_llm_start(inputs, **self._identifying_params()) 14 15 payload = { 16 "model": self.model_name, 17 "messages": [{"role": m.type, "content": m.content} for m in messages], 18 "temperature": self.temperature, 19 "max_tokens": self.max_tokens, 20 **({"stop": stop} if stop else {}), 21 } 22 23 async with httpx.AsyncClient(timeout=self.timeout) as client: 24 for attempt in range(self.max_retries): 25 try: 26 resp = await client.post( 27 self.endpoint_url, 28 json=payload, 29 headers={"Authorization": f"Bearer {self.api_key}"}, 30 ) 31 resp.raise_for_status() 32 data = resp.json() 33 break 34 except Exception: 35 if attempt + 1 == self.max_retries: 36 raise 37 await asyncio.sleep(2 ** attempt) 38 39 text = data["choices"][0]["message"]["content"] 40 usage = data.get("usage", {}) 41 42 generation = ChatGeneration( 43 message=AIMessage(content=text), 44 generation_info={"model": self.model_name}, 45 usage=usage, 46 ) 47 await run_manager.on_llm_end([generation], **self._identifying_params()) 48 return ChatResult(generations=[generation]) 49

2.4 _identifying_params

Python
1 def _identifying_params(self) -> Dict[str, Any]: 2 return { 3 "model": self.model_name, 4 "temperature": self.temperature, 5 "max_tokens": self.max_tokens, 6 } 7

This ensures LangChain can key cache and telemetry by unique model settings.


3. Streaming & Callback Integration

When your API offers streaming (chunked) responses, you can surface partial tokens to LangChain’s streaming interface.

3.1 Streaming Method Signature

Python
1from typing import Iterator 2from langchain_core.outputs import ChatGenerationChunk 3from langchain_core.messages import AIMessageChunk 4 5 def _stream( 6 self, 7 messages: List[BaseMessage], 8 stop: Optional[List[str]] = None, 9 run_manager: Optional[CallbackManagerForLLMRun] = None, 10 **kwargs, 11 ) -> Iterator[ChatGenerationChunk]: 12 # 1. Start event 13 inputs = [m.content for m in messages] 14 run_manager.on_llm_start(inputs, **self._identifying_params()) 15 16 # 2. Initiate streaming request 17 resp = requests.post( 18 self.endpoint_url, 19 json={…}, 20 headers={…}, 21 stream=True, 22 timeout=self.timeout, 23 ) 24 25 # 3. Iterate over chunks 26 for chunk in resp.iter_lines(): 27 if not chunk: 28 continue 29 part = json.loads(chunk.decode("utf-8")) 30 token = part["choices"][0]["delta"].get("content", "") 31 # Emit stream event 32 run_manager.on_llm_stream(token) 33 yield ChatGenerationChunk( 34 message=AIMessageChunk(content=token), 35 generation_info={"model": self.model_name}, 36 ) 37 38 # 4. End event 39 run_manager.on_llm_end([], **self._identifying_params()) 40

3.2 Async Streaming

With httpx.AsyncClient:

Python
1 async def _astream( 2 self, 3 messages, stop=None, run_manager=None, **kwargs 4 ) -> AsyncIterator[ChatGenerationChunk]: 5 inputs = [m.content for m in messages] 6 await run_manager.on_llm_start(inputs, **self._identifying_params()) 7 8 async with httpx.AsyncClient(timeout=self.timeout) as client: 9 async with client.stream( 10 "POST", self.endpoint_url, json={…}, headers={…} 11 ) as resp: 12 async for line in resp.aiter_lines(): 13 if not line: 14 continue 15 part = json.loads(line) 16 token = part["choices"][0]["delta"].get("content", "") 17 await run_manager.on_llm_stream(token) 18 yield ChatGenerationChunk( 19 message=AIMessageChunk(content=token), 20 generation_info={"model": self.model_name}, 21 ) 22 await run_manager.on_llm_end([], **self._identifying_params()) 23

3.3 Astream Events API

LangChain’s newer astream_events lets you attach a single event generator:

Python
1 def astream_events(self, messages, **kwargs): 2 async def event_generator(): 3 # mix of stream and callbacks 4 … 5 return event_generator() 6

Use this if you want unified async event handling and avoid duplicating callback calls.


4. Testing, Batching, Error Handling, Deployment

A production-ready custom chat model also includes:

4.1 Batch Support (Threadpool)

LangChain auto-wraps your sync _generate into batch calls via a threadpool by default. To manually customize:

Python
1 def batch( 2 self, 3 prompts: List[List[BaseMessage]], 4 **kwargs 5 ) -> List[ChatResult]: 6 # prompt-level batching logic, e.g., multi-prompt HTTP call 7 … 8

Or rely on inherited .batch() behavior.

4.2 Unit & Integration Tests

  • Sync invoke:
    Python
    1model = AcmeChatModel(api_key="x", model="acme-chat") 2res = model.invoke([HumanMessage(content="Hello")]) 3assert isinstance(res, ChatResult) 4assert res.generations[0].message.content.startswith("…") 5
  • Async ainvoke:
    Python
    1import asyncio 2async def test_async(): 3 model = AcmeChatModel(…) 4 res = await model.ainvoke([HumanMessage("Hi")]) 5 assert … 6asyncio.run(test_async()) 7
  • Streaming: iterate and assemble chunks.

4.3 Error Handling & Retries

  • Use exponential backoff:
    Python
    1for i in range(self.max_retries): 2 try: …; break 3 except requests.RequestException: 4 time.sleep(2 ** i) 5
  • Catch HTTP 5xx and 429 codes specially.
  • Raise on final failure.

4.4 Deployment & Configuration

  • Read secrets from environment:
    Python
    1import os 2model = AcmeChatModel( 3 api_key=os.getenv("ACME_API_KEY"), 4 model="acme-chat", 5) 6

Recommended Articles

Discover more articles you might find interesting

Implementing LangGraph REST API with FastAPI
Technical Insights

Implementing LangGraph REST API with FastAPI

This guide provides a comprehensive implementation plan for building a LangGraph REST API using FastAPI, covering environment setup, agent definitions, endpoint creation, testing, and deployment.

2101050
Jun 18
149
Read More
DeepSite v2 Practical Guide
Technical Insights

DeepSite v2 Practical Guide

A comprehensive guide to DeepSite v2, covering its features, installation, and advanced workflows.

2101050
Jun 21
110
Read More
Fastify OpenTelemetry: Logging, Metrics, and Tracing in Practice
Technical Insights

Fastify OpenTelemetry: Logging, Metrics, and Tracing in Practice

Learn how to implement logging, metrics, and tracing in Fastify using OpenTelemetry.

2101050
Jul 11
105
Read More
Creating Diverse Logo Designs with Flux Model and ComfyUI
Technical Insights

Creating Diverse Logo Designs with Flux Model and ComfyUI

Learn to leverage the Flux model and ComfyUI for unique logo designs through effective prompts and examples.

2101050
Jan 10
92
Read More
Formatting Dates in TypeScript to UTC
Technical Insights

Formatting Dates in TypeScript to UTC

A guide on how to format dates in TypeScript to the specific format YYYY-MM-DDTHH:mm:ss+00:00.

2101050
Dec 19
82
Read More
EAS Local Build Expo for Windows: Step-by-Step Guide, Troubleshooting, and Real-World Cases
Technical Insights

EAS Local Build Expo for Windows: Step-by-Step Guide, Troubleshooting, and Real-World Cases

Explore the complete process of setting up and troubleshooting EAS Local Builds on Windows using WSL, including practical examples and best practices.

2101050
Jul 08
77
Read More