codekingpro/portable-devtools
114k
1"""Tool for the Nuclia Understanding API.2 3Installation:4 5```bash6 pip install --upgrade protobuf7 pip install nucliadb-protos8```9"""10 11import asyncio12import base6413import logging14import mimetypes15import os16from typing import Any, Dict, Optional, Type, Union17 18import requests19from langchain_core.callbacks import (20 AsyncCallbackManagerForToolRun,21 CallbackManagerForToolRun,22)23from langchain_core.tools import BaseTool24from pydantic import BaseModel, Field25 26logger = logging.getLogger(__name__)27 28 29class NUASchema(BaseModel):30 """Input for Nuclia Understanding API.31 32 Attributes:33 action: Action to perform. Either `push` or `pull`.34 id: ID of the file to push or pull.35 path: Path to the file to push (needed only for `push` action).36 text: Text content to process (needed only for `push` action).37 """38 39 action: str = Field(40 ...,41 description="Action to perform. Either `push` or `pull`.",42 )43 id: str = Field(44 ...,45 description="ID of the file to push or pull.",46 )47 path: Optional[str] = Field(48 ...,49 description="Path to the file to push (needed only for `push` action).",50 )51 text: Optional[str] = Field(52 ...,53 description="Text content to process (needed only for `push` action).",54 )55 56 57class NucliaUnderstandingAPI(BaseTool):58 """Tool to process files with the Nuclia Understanding API."""59 60 name: str = "nuclia_understanding_api"61 description: str = (62 "A wrapper around Nuclia Understanding API endpoints. "63 "Useful for when you need to extract text from any kind of files. "64 )65 args_schema: Type[BaseModel] = NUASchema66 _results: Dict[str, Any] = {}67 _config: Dict[str, Any] = {}68 69 def __init__(self, enable_ml: bool = False) -> None:70 zone = os.environ.get("NUCLIA_ZONE", "europe-1")71 self._config["BACKEND"] = f"https://{zone}.nuclia.cloud/api/v1"72 key = os.environ.get("NUCLIA_NUA_KEY")73 if not key:74 raise ValueError("NUCLIA_NUA_KEY environment variable not set")75 else:76 self._config["NUA_KEY"] = key77 self._config["enable_ml"] = enable_ml78 super().__init__()79 80 def _run(81 self,82 action: str,83 id: str,84 path: Optional[str],85 text: Optional[str],86 run_manager: Optional[CallbackManagerForToolRun] = None,87 ) -> str:88 """Use the tool."""89 if action == "push":90 self._check_params(path, text)91 if path:92 return self._pushFile(id, path)93 if text:94 return self._pushText(id, text)95 elif action == "pull":96 return self._pull(id)97 return ""98 99 async def _arun(100 self,101 action: str,102 id: str,103 path: Optional[str] = None,104 text: Optional[str] = None,105 run_manager: Optional[AsyncCallbackManagerForToolRun] = None,106 ) -> str:107 """Use the tool asynchronously."""108 self._check_params(path, text)109 if path:110 self._pushFile(id, path)111 if text:112 self._pushText(id, text)113 data = None114 while True:115 data = self._pull(id)116 if data:117 break118 await asyncio.sleep(15)119 return data120 121 def _pushText(self, id: str, text: str) -> str:122 field = {123 "textfield": {"text": {"body": text, "format": 0}},124 "processing_options": {"ml_text": self._config["enable_ml"]},125 }126 return self._pushField(id, field)127 128 def _pushFile(self, id: str, content_path: str) -> str:129 with open(content_path, "rb") as source_file:130 response = requests.post(131 self._config["BACKEND"] + "/processing/upload",132 headers={133 "content-type": mimetypes.guess_type(content_path)[0]134 or "application/octet-stream",135 "x-stf-nuakey": "Bearer " + self._config["NUA_KEY"],136 },137 data=source_file.read(),138 )139 if response.status_code != 200:140 logger.info(141 f"Error uploading {content_path}: "142 f"{response.status_code} {response.text}"143 )144 return ""145 else:146 field = {147 "filefield": {"file": f"{response.text}"},148 "processing_options": {"ml_text": self._config["enable_ml"]},149 }150 return self._pushField(id, field)151 152 def _pushField(self, id: str, field: Any) -> str:153 logger.info(f"Pushing {id} in queue")154 response = requests.post(155 self._config["BACKEND"] + "/processing/push",156 headers={157 "content-type": "application/json",158 "x-stf-nuakey": "Bearer " + self._config["NUA_KEY"],159 },160 json=field,161 )162 if response.status_code != 200:163 logger.info(164 f"Error pushing field {id}:{response.status_code} {response.text}"165 )166 raise ValueError("Error pushing field")167 else:168 uuid = response.json()["uuid"]169 logger.info(f"Field {id} pushed in queue, uuid: {uuid}")170 self._results[id] = {"uuid": uuid, "status": "pending"}171 return uuid172 173 def _pull(self, id: str) -> str:174 self._pull_queue()175 result = self._results.get(id, None)176 if not result:177 logger.info(f"{id} not in queue")178 return ""179 elif result["status"] == "pending":180 logger.info(f"Waiting for {result['uuid']} to be processed")181 return ""182 else:183 return result["data"]184 185 def _pull_queue(self) -> None:186 try:187 from nucliadb_protos.writer_pb2 import BrokerMessage188 except ImportError as e:189 raise ImportError(190 "nucliadb-protos is not installed. "191 "Run `pip install nucliadb-protos` to install."192 ) from e193 try:194 from google.protobuf.json_format import MessageToJson195 except ImportError as e:196 raise ImportError(197 "Unable to import google.protobuf, please install with "198 "`pip install protobuf`."199 ) from e200 201 res = requests.get(202 self._config["BACKEND"] + "/processing/pull",203 headers={204 "x-stf-nuakey": "Bearer " + self._config["NUA_KEY"],205 },206 ).json()207 if res["status"] == "empty":208 logger.info("Queue empty")209 elif res["status"] == "ok":210 payload = res["payload"]211 pb = BrokerMessage()212 pb.ParseFromString(base64.b64decode(payload))213 uuid = pb.uuid214 logger.info(f"Pulled {uuid} from queue")215 matching_id = self._find_matching_id(uuid)216 if not matching_id:217 logger.info(f"No matching id for {uuid}")218 else:219 self._results[matching_id]["status"] = "done"220 data = MessageToJson( # type: ignore[call-arg]221 pb,222 preserving_proto_field_name=True,223 including_default_value_fields=True,224 )225 self._results[matching_id]["data"] = data226 227 def _find_matching_id(self, uuid: str) -> Union[str, None]:228 for id, result in self._results.items():229 if result["uuid"] == uuid:230 return id231 return None232 233 def _check_params(self, path: Optional[str], text: Optional[str]) -> None:234 if not path and not text:235 raise ValueError("File path or text is required")236 if path and text:237 raise ValueError("Cannot process both file and text on a single run")238 