mirror of
https://github.com/Mintplex-Labs/langchain-python.git
synced 2026-07-21 00:35:23 -04:00
6e90406e0f
<!-- Thank you for contributing to LangChain! Your PR will appear in our release under the title you set. Please make sure it highlights your valuable contribution. Replace this with a description of the change, the issue it fixes (if applicable), and relevant context. List any dependencies required for this change. After you're done, someone will review your PR. They may suggest improvements. If no one reviews your PR within a few days, feel free to @-mention the same people again, as notifications can get lost. Finally, we'd love to show appreciation for your contribution - if you'd like us to shout you out on Twitter, please also include your handle! --> I used the APIChain sometimes it failed during the intermediate step when generating the api url and calling the `request` function. After some digging, I found the url sometimes includes the space at the beginning, like `%20https://...api.com` which causes the ` self.requests_wrapper.get` internal function to fail. Including a little string preprocessing `.strip` to remove the space seems to improve the robustness of the APIchain to make sure it can send the request and retrieve the API result more reliably. <!-- Remove if not applicable --> Fixes # (issue) #### Before submitting <!-- If you're adding a new integration, please include: 1. a test for the integration - favor unit tests that does not rely on network access. 2. an example notebook showing its use See contribution guidelines for more information on how to write tests, lint etc: https://github.com/hwchase17/langchain/blob/master/.github/CONTRIBUTING.md --> #### Who can review? @vowelparrot Tag maintainers/contributors who might be interested: <!-- For a quicker response, figure out the right person to tag with @ @hwchase17 - project lead Tracing / Callbacks - @agola11 Async - @agola11 DataLoaders - @eyurtsev Models - @hwchase17 - @agola11 Agents / Tools / Toolkits - @vowelparrot VectorStores / Retrievers / Memory - @dev2049 -->
149 lines
5.2 KiB
Python
149 lines
5.2 KiB
Python
"""Chain that makes API calls and summarizes the responses to answer a question."""
|
|
from __future__ import annotations
|
|
|
|
from typing import Any, Dict, List, Optional
|
|
|
|
from pydantic import Field, root_validator
|
|
|
|
from langchain.base_language import BaseLanguageModel
|
|
from langchain.callbacks.manager import (
|
|
AsyncCallbackManagerForChainRun,
|
|
CallbackManagerForChainRun,
|
|
)
|
|
from langchain.chains.api.prompt import API_RESPONSE_PROMPT, API_URL_PROMPT
|
|
from langchain.chains.base import Chain
|
|
from langchain.chains.llm import LLMChain
|
|
from langchain.prompts import BasePromptTemplate
|
|
from langchain.requests import TextRequestsWrapper
|
|
|
|
|
|
class APIChain(Chain):
|
|
"""Chain that makes API calls and summarizes the responses to answer a question."""
|
|
|
|
api_request_chain: LLMChain
|
|
api_answer_chain: LLMChain
|
|
requests_wrapper: TextRequestsWrapper = Field(exclude=True)
|
|
api_docs: str
|
|
question_key: str = "question" #: :meta private:
|
|
output_key: str = "output" #: :meta private:
|
|
|
|
@property
|
|
def input_keys(self) -> List[str]:
|
|
"""Expect input key.
|
|
|
|
:meta private:
|
|
"""
|
|
return [self.question_key]
|
|
|
|
@property
|
|
def output_keys(self) -> List[str]:
|
|
"""Expect output key.
|
|
|
|
:meta private:
|
|
"""
|
|
return [self.output_key]
|
|
|
|
@root_validator(pre=True)
|
|
def validate_api_request_prompt(cls, values: Dict) -> Dict:
|
|
"""Check that api request prompt expects the right variables."""
|
|
input_vars = values["api_request_chain"].prompt.input_variables
|
|
expected_vars = {"question", "api_docs"}
|
|
if set(input_vars) != expected_vars:
|
|
raise ValueError(
|
|
f"Input variables should be {expected_vars}, got {input_vars}"
|
|
)
|
|
return values
|
|
|
|
@root_validator(pre=True)
|
|
def validate_api_answer_prompt(cls, values: Dict) -> Dict:
|
|
"""Check that api answer prompt expects the right variables."""
|
|
input_vars = values["api_answer_chain"].prompt.input_variables
|
|
expected_vars = {"question", "api_docs", "api_url", "api_response"}
|
|
if set(input_vars) != expected_vars:
|
|
raise ValueError(
|
|
f"Input variables should be {expected_vars}, got {input_vars}"
|
|
)
|
|
return values
|
|
|
|
def _call(
|
|
self,
|
|
inputs: Dict[str, Any],
|
|
run_manager: Optional[CallbackManagerForChainRun] = None,
|
|
) -> Dict[str, str]:
|
|
_run_manager = run_manager or CallbackManagerForChainRun.get_noop_manager()
|
|
question = inputs[self.question_key]
|
|
api_url = self.api_request_chain.predict(
|
|
question=question,
|
|
api_docs=self.api_docs,
|
|
callbacks=_run_manager.get_child(),
|
|
)
|
|
_run_manager.on_text(api_url, color="green", end="\n", verbose=self.verbose)
|
|
api_url = api_url.strip()
|
|
api_response = self.requests_wrapper.get(api_url)
|
|
_run_manager.on_text(
|
|
api_response, color="yellow", end="\n", verbose=self.verbose
|
|
)
|
|
answer = self.api_answer_chain.predict(
|
|
question=question,
|
|
api_docs=self.api_docs,
|
|
api_url=api_url,
|
|
api_response=api_response,
|
|
callbacks=_run_manager.get_child(),
|
|
)
|
|
return {self.output_key: answer}
|
|
|
|
async def _acall(
|
|
self,
|
|
inputs: Dict[str, Any],
|
|
run_manager: Optional[AsyncCallbackManagerForChainRun] = None,
|
|
) -> Dict[str, str]:
|
|
_run_manager = run_manager or AsyncCallbackManagerForChainRun.get_noop_manager()
|
|
question = inputs[self.question_key]
|
|
api_url = await self.api_request_chain.apredict(
|
|
question=question,
|
|
api_docs=self.api_docs,
|
|
callbacks=_run_manager.get_child(),
|
|
)
|
|
await _run_manager.on_text(
|
|
api_url, color="green", end="\n", verbose=self.verbose
|
|
)
|
|
api_url = api_url.strip()
|
|
api_response = await self.requests_wrapper.aget(api_url)
|
|
await _run_manager.on_text(
|
|
api_response, color="yellow", end="\n", verbose=self.verbose
|
|
)
|
|
answer = await self.api_answer_chain.apredict(
|
|
question=question,
|
|
api_docs=self.api_docs,
|
|
api_url=api_url,
|
|
api_response=api_response,
|
|
callbacks=_run_manager.get_child(),
|
|
)
|
|
return {self.output_key: answer}
|
|
|
|
@classmethod
|
|
def from_llm_and_api_docs(
|
|
cls,
|
|
llm: BaseLanguageModel,
|
|
api_docs: str,
|
|
headers: Optional[dict] = None,
|
|
api_url_prompt: BasePromptTemplate = API_URL_PROMPT,
|
|
api_response_prompt: BasePromptTemplate = API_RESPONSE_PROMPT,
|
|
**kwargs: Any,
|
|
) -> APIChain:
|
|
"""Load chain from just an LLM and the api docs."""
|
|
get_request_chain = LLMChain(llm=llm, prompt=api_url_prompt)
|
|
requests_wrapper = TextRequestsWrapper(headers=headers)
|
|
get_answer_chain = LLMChain(llm=llm, prompt=api_response_prompt)
|
|
return cls(
|
|
api_request_chain=get_request_chain,
|
|
api_answer_chain=get_answer_chain,
|
|
requests_wrapper=requests_wrapper,
|
|
api_docs=api_docs,
|
|
**kwargs,
|
|
)
|
|
|
|
@property
|
|
def _chain_type(self) -> str:
|
|
return "api_chain"
|