Open WebUI RAG Filter

(compatible with Retrieval Suite Version 1.14.5 onwards)

"""
title: Enterprise Knowledge Filter
author: RheinInsights
author_url:
version: 1.0.2
required_open_webui_version: 0.3.30
"""

import asyncio
import time
from typing import Optional, Callable, Awaitable
from pydantic import BaseModel, Field
import requests
import urllib3
import logging

urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning)


class Filter:
    class Valves(BaseModel):
        RheininsightsUrl: str = Field(default="https://localhost")
        AuthenticationToken: str = Field(default="secret")
        QueryPipelineId: int = Field(default=0)
        priority: int = Field(default=0)
        enabled: bool = Field(default=True)

    def __init__(self):
        self.valves = self.Valves()
        self.logger = logging.getLogger("api_usage")

    async def inlet(
        self,
        body: dict,
        __user__: Optional[dict] = None,
        __event_emitter__: Optional[Callable[[dict], Awaitable[dict]]] = None,
        __metadata__: Optional[dict] = None,
    ) -> dict:

        if not self.valves.enabled:
            return body

        user_email = __user__.get("email", "unknown") if __user__ else "anonymous"
        model = body.get("model", "unknown")
        interface = __metadata__.get("interface", "api") if __metadata__ else "api"
        chat_id = __metadata__.get("chat_id") if __metadata__ else None

        print(
            f"Request: user={user_email}, model={model}, "
            f"interface={interface}, chat_id={chat_id or 'none'}"
        )

        user_input = self._extract_user_input(body)
        if not user_input:
            return body

        context = body["messages"]

        user_mail = __user__.get("email")

        do_not_post_updates = "### Task:" in user_input

        try:
            knowledge = await self._search_knowledge(
                query=user_input,
                context=context,
                mail=user_mail,
                __event_emitter__=__event_emitter__,
                do_not_post_updates=do_not_post_updates,
            )

            await self.post_final_message(__event_emitter__, do_not_post_updates)

            if not knowledge:
                return body

            body["messages"] = self._inject_knowledge(context, knowledge)
            return body

        except Exception as e:
            await self.emit_event(
                f"Knowledge filter error: {str(e)}",
                __event_emitter__,
                do_not_post_updates,
                done=True,
            )
            return body

    async def outlet(
        self,
        body: dict,
        __user__: Optional[dict] = None,
        __event_emitter__: Optional[Callable[[dict], Awaitable[dict]]] = None,
        __metadata__: Optional[dict] = None,
    ) -> dict:
        return body

    def _extract_user_input(self, body: dict) -> str:
        messages = body.get("messages", [])
        if not isinstance(messages, list) or not messages:
            return ""

        last_message = messages[-1]
        content = last_message.get("content", "")

        if isinstance(content, list):
            for item in content:
                if item.get("type") == "text":
                    return item.get("text", "")
            return ""

        return content or ""

    def _inject_knowledge(self, messages: list, knowledge: str) -> list:
        if not knowledge.strip():
            return messages

        system_message = {
            "role": "system",
            "content": (
                "Use the following enterprise knowledge if relevant:\n\n" f"{knowledge}"
            ),
        }

        return [system_message, *messages]

    async def _search_knowledge(
        self,
        query: str,
        context: object,
        mail: str,
        __event_emitter__: Optional[Callable[[dict], Awaitable[dict]]],
        do_not_post_updates: bool,
    ) -> Optional[str]:
        data = await asyncio.to_thread(self.submit_query, query, context, mail)

        if not data or not data.get("threadId"):
            return None

        result = await self.wait_for_finalization(
            data["threadId"], mail, __event_emitter__, do_not_post_updates
        )

        if not result or result.get("isFailed"):
            return None

        return self.handle_results(result.get("result", {}))

    async def wait_for_finalization(
        self,
        thread_id: str,
        mail: str,
        __event_emitter__: Optional[Callable[[dict], Awaitable[dict]]],
        do_not_post_updates: bool,
    ) -> Optional[dict]:
        deadline = time.time() + 600
        interval = 1
        last_update_at = 0

        await self.emit_event("Processing", __event_emitter__, do_not_post_updates)

        while time.time() < deadline:
            await asyncio.sleep(interval)
            data = await asyncio.to_thread(self.get_status, thread_id, mail)

            if not data:
                return None

            if data.get("status") == "UNKNOWN":
                return None

            if data.get("status") == "FINISHED":
                return data

            msg = data.get("message")
            if not msg:
                continue

            if msg.get("lastUpdate") == last_update_at:
                continue

            await self.emit_event(
                msg.get("message", "Processing"),
                __event_emitter__,
                do_not_post_updates,
            )
            last_update_at = msg.get("lastUpdate", 0)

        return None

    def submit_query(self, query: str, context: object, mail: str) -> dict:
        url = (
            f"{self.valves.RheininsightsUrl}"
            f"/api/v1/querypipelines/search/async?queryPipelineId={self.valves.QueryPipelineId}"
        )
        headers = {"Authorization": self.valves.AuthenticationToken}

        query_object = {
            "query": query,
            "userPrincipalName": mail,
            "context": context,
        }

        response = requests.post(
            url, headers=headers, json=query_object, verify=False, timeout=120
        )
        response.raise_for_status()
        return response.json()

    def get_status(self, thread_id: str, mail: str) -> dict:
        url = (
            f"{self.valves.RheininsightsUrl}/api/v1/querypipelines/search/async/status"
        )
        headers = {"Authorization": self.valves.AuthenticationToken}

        query_object = {"threadId": thread_id, "userPrincipalName": mail}

        response = requests.post(
            url, headers=headers, json=query_object, verify=False, timeout=120
        )
        response.raise_for_status()
        return response.json()

    async def post_final_message(
        self,
        __event_emitter__: Optional[Callable[[dict], Awaitable[dict]]],
        do_not_post_updates: bool,
    ):
        if not __event_emitter__ or do_not_post_updates:
            return

        await __event_emitter__(
            {
                "type": "status",
                "data": {"description": "Completed successfully", "done": True},
            }
        )

    async def emit_event(
        self,
        msg: str,
        __event_emitter__: Optional[Callable[[dict], Awaitable[dict]]],
        do_not_post_updates: bool,
    ):
        if not __event_emitter__ or do_not_post_updates:
            return

        await __event_emitter__(
            {
                "type": "status",
                "data": {
                    "status": "in_progress",
                    "description": msg,
                    "done": False,
                },
            }
        )

    def handle_results(self, data: dict) -> str:
        try:
            return "\n\n".join(
                r.get("teaser", "") for r in data.get("results", []) if r.get("teaser")
            )
        except Exception:
            return ""