Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tool.py238 linesDownload Raw Back to nuclia
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 
codekingpro/portable-devtools · Team Ai