Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions docs/api/deployment.md
Original file line number Diff line number Diff line change
Expand Up @@ -26,15 +26,15 @@ make run
Or using CLI:

```bash
llm_guard_api ./config/scanners.yml
llm_guard_api ./config/app_config.yml
```

### Using gunicorn

In case you want to use `gunicorn` to run the API, you can use the following command:

```bash
gunicorn --workers 1 --preload --worker-class uvicorn.workers.UvicornWorker 'app.app:create_app(config_file="./config/scanners.yml")'
gunicorn --workers 1 --preload --worker-class uvicorn.workers.UvicornWorker 'app.app:create_app(config_file="./config/app_config.yml")'
```

It will preload models in the shared memory among workers, which can be useful for performance.
Expand Down Expand Up @@ -67,7 +67,7 @@ This will start the API on port 8000. You can now access the API at `http://loca
If you want to use a custom configuration, you can mount a volume to `/home/user/app/config`:

```bash
docker run -d -p 8000:8000 -e APP_WORKERS=1 -e AUTH_TOKEN='my-token' -e LOG_LEVEL='DEBUG' -v ./entrypoint.sh:/home/user/app/entrypoint.sh -v ./config/scanners.yml:/home/user/app/config/scanners.yml laiyer/llm-guard-api:latest
docker run -d -p 8000:8000 -e APP_WORKERS=1 -e AUTH_TOKEN='my-token' -e LOG_LEVEL='DEBUG' -v ./entrypoint.sh:/home/user/app/entrypoint.sh -v ./config/app_config.yml:/home/user/app/config/app_config.yml laiyer/llm-guard-api:latest
```

!!! warning
Expand Down
2 changes: 1 addition & 1 deletion tueri/input_scanners/anonymize.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,7 @@ def scan(self, prompt: str) -> tuple[str, bool, float]:
self._vault.append((placeholder, original_value))
return (
self._preamble + sanitized_prompt,
False,
True,
calculate_risk_score(risk_score, self._threshold),
)

Expand Down
2 changes: 1 addition & 1 deletion tueri_api/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ RUN pip install --no-cache-dir --upgrade pip && \

RUN python -m spacy download en_core_web_sm

COPY --chown=user:user tueri_api/config/scanners.yml ./config/scanners.yml
COPY --chown=user:user tueri_api/config/app_config.yml ./config/app_config.yml
COPY --chown=user:user tueri_api/entrypoint.sh ./entrypoint.sh

RUN chmod +x ./entrypoint.sh
Expand Down
2 changes: 1 addition & 1 deletion tueri_api/Dockerfile-cuda
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,7 @@ RUN pip3 install --no-cache-dir --upgrade pip && \

RUN python -m spacy download en_core_web_sm

COPY --chown=user:user ./config/scanners.yml ./config/scanners.yml
COPY --chown=user:user ./config/app_config.yml ./config/app_config.yml
COPY --chown=user:user entrypoint.sh ./entrypoint.sh

RUN chmod +x ./entrypoint.sh
Expand Down
18 changes: 9 additions & 9 deletions tueri_api/app/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,7 @@


def create_app() -> FastAPI:
config_file = os.getenv("CONFIG_FILE", "./config/scanners.yml")
config_file = os.getenv("CONFIG_FILE", "./config/app_config.yml")
if not config_file:
raise ValueError("Config file is required")

Expand Down Expand Up @@ -124,15 +124,15 @@ async def check_auth(credentials: credentials_type) -> bool:
def _get_input_scanners_function(config: Config, vault: Vault) -> Callable:
scanners = []
if not config.app.lazy_load:
LOGGER.debug("Loading input scanners")
scanners = get_input_scanners(config.input_scanners, vault)
LOGGER.debug("Loading input scanners from MongoDB")
scanners = get_input_scanners([], vault)

def get_cached_scanners() -> List[InputScanner]:
nonlocal scanners

if not scanners and config.app.lazy_load:
LOGGER.debug("Lazy loading input scanners")
scanners = get_input_scanners(config.input_scanners, vault)
LOGGER.debug("Lazy loading input scanners from MongoDB")
scanners = get_input_scanners([], vault)

return scanners

Expand All @@ -142,15 +142,15 @@ def get_cached_scanners() -> List[InputScanner]:
def _get_output_scanners_function(config: Config, vault: Vault) -> Callable:
scanners = []
if not config.app.lazy_load:
LOGGER.debug("Loading output scanners")
scanners = get_output_scanners(config.output_scanners, vault)
LOGGER.debug("Loading output scanners from MongoDB")
scanners = get_output_scanners([], vault)

def get_cached_scanners() -> List[OutputScanner]:
nonlocal scanners

if not scanners and config.app.lazy_load:
LOGGER.debug("Lazy loading output scanners")
scanners = get_output_scanners(config.output_scanners, vault)
LOGGER.debug("Lazy loading output scanners from MongoDB")
scanners = get_output_scanners([], vault)

return scanners

Expand Down
4 changes: 2 additions & 2 deletions tueri_api/app/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,8 +50,8 @@ class ScannerConfig(BaseModel):


class Config(BaseModel):
input_scanners: List[ScannerConfig] = Field()
output_scanners: List[ScannerConfig] = Field()
input_scanners: List[ScannerConfig] = Field(default_factory=list)
output_scanners: List[ScannerConfig] = Field(default_factory=list)
rate_limit: RateLimitConfig = Field(default_factory=RateLimitConfig)
auth: Optional[AuthConfig] = Field(default=None)
app: AppConfig = Field(default_factory=AppConfig)
Expand Down
54 changes: 41 additions & 13 deletions tueri_api/app/scanner.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
import asyncio
import time
import os
import logging
from typing import Dict, List, Optional

import structlog
import torch
from opentelemetry import metrics
from pymongo import MongoClient

from tueri import input_scanners, output_scanners
from tueri.input_scanners.ban_competitors import MODEL_V1 as BAN_COMPETITORS_MODEL
Expand All @@ -27,52 +30,77 @@

LOGGER = structlog.getLogger(__name__)

MONGO_URI = os.getenv("MONGO_URI", "mongodb://root:example@localhost:27017/")
MONGO_DB, MONGO_COLLECTION = os.getenv("MONGO_DB", "ChatApp"), os.getenv("MONGO_COLLECTION", "TueriScanners")

# Suppress MongoDB heartbeat logs
logging.getLogger("pymongo.topology").setLevel(logging.WARNING)
logging.getLogger("pymongo.serverSelection").setLevel(logging.WARNING)

try:
mongo_client = MongoClient(MONGO_URI, serverSelectionTimeoutMS=5000, heartbeatFrequencyMS=60000)
db = mongo_client[MONGO_DB]
scanners_collection = db[MONGO_COLLECTION]
mongo_client.admin.command("ping")
except Exception as e:
LOGGER.error("Error connecting to MongoDB", error=str(e))
raise

meter = metrics.get_meter_provider().get_meter(__name__)
scanners_valid_counter = meter.create_counter(
name="scanners.valid",
unit="1",
description="measures the number of valid scanners",
)

def _fetch_scanners_from_mongo(scanner_type: str) -> List[ScannerConfig]:
coll = scanners_collection.find({"type": scanner_type})
scanners: List[ScannerConfig] = []
for scanner in coll:
scanners.append(ScannerConfig(
type=scanner.get("id"),
params=scanner.get("params", {})))
return scanners

def get_input_scanners(scanners: List[ScannerConfig], vault: Vault) -> List[InputScanner]:
def get_input_scanners(scanners: List[ScannerConfig], vault: Vault) -> List[InputScanner]:
"""
Load input scanners from the configuration file.
Load input scanners from MongoDB.
"""

input_scanners_loaded = []
for scanner in scanners:
input_scanners_config = _fetch_scanners_from_mongo("input")
loaded_input_scanners: List[InputScanner] = []
for scanner in input_scanners_config:
LOGGER.debug("Loading input scanner", scanner=scanner.type, **get_resource_utilization())
input_scanners_loaded.append(
loaded_input_scanners.append(
_get_input_scanner(
scanner.type,
scanner.params,
vault=vault,
)
)

return input_scanners_loaded
return loaded_input_scanners


def get_output_scanners(scanners: List[ScannerConfig], vault: Vault) -> List[OutputScanner]:
"""
Load output scanners from the configuration file.
Load output scanners from MongoDB.
"""
output_scanners_loaded = []
for scanner in scanners:
output_scanners_config = _fetch_scanners_from_mongo("output")
loaded_output_scanners: List[OutputScanner] = []
for scanner in output_scanners_config:
LOGGER.debug("Loading output scanner", scanner=scanner.type, **get_resource_utilization())
output_scanners_loaded.append(
loaded_output_scanners.append(
_get_output_scanner(
scanner.type,
scanner.params,
vault=vault,
)
)

return output_scanners_loaded
return loaded_output_scanners


def _configure_model(model: Model, scanner_config: Optional[Dict]):
def _configure_model(model: Model, scanner_config: Optional[Dict]):
if scanner_config is None:
scanner_config = {}

Expand Down
Loading