Coverage for open_webui/retrieval/external.py: 9%
182 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1import asyncio
2import logging
3import re
4import time
5from typing import Any, Optional
7from open_webui.config import RAG_EMBEDDING_QUERY_PREFIX
8from open_webui.models.config import Config
9from open_webui.models.knowledge import KnowledgeModel
11log = logging.getLogger(__name__)
13EXTERNAL_KNOWLEDGE_CONNECTIONS_CONFIG_KEY = 'external_knowledge.connections'
14IDENTIFIER_RE = re.compile(r'^[A-Za-z_][A-Za-z0-9_]*$')
17async def _get_external_connection(connection_id: str) -> Optional[dict]:
18 connections = await Config.get(EXTERNAL_KNOWLEDGE_CONNECTIONS_CONFIG_KEY, []) or []
19 return next((connection for connection in connections if connection.get('id') == connection_id), None)
22def _get_path(data: Any, path: Optional[str], default=None):
23 if not path:
24 return default
25 value = data
26 for part in path.split('.'):
27 if isinstance(value, dict):
28 value = value.get(part, default)
29 else:
30 return default
31 return value
34def _normalize_result(result: dict, mapping: dict, knowledge: KnowledgeModel, distance: Optional[float] = None) -> dict:
35 content = _get_path(result, mapping.get('content_field', 'content'), '')
36 title = _get_path(result, mapping.get('title_field', 'title'), None)
37 source = _get_path(result, mapping.get('source_field', 'source'), None)
38 url = _get_path(result, mapping.get('url_field', 'url'), None)
39 document_id = _get_path(result, mapping.get('document_id_field', 'document_id'), None)
40 page = _get_path(result, mapping.get('page_field', 'page'), None)
41 metadata = _get_path(result, mapping.get('metadata_field', 'metadata'), {}) or {}
42 score = _get_path(result, mapping.get('score_field', 'score'), distance)
44 if not isinstance(metadata, dict):
45 metadata = {'external_metadata': metadata}
47 source_name = source or title or metadata.get('source') or metadata.get('name') or knowledge.name
48 metadata.update(
49 {
50 'name': title or source_name,
51 'source': source_name,
52 'url': url,
53 'file_id': document_id or f'external-{knowledge.id}',
54 'knowledge_id': knowledge.id,
55 'knowledge_name': knowledge.name,
56 'external': True,
57 }
58 )
59 if page is not None:
60 metadata['page'] = page
61 if document_id is not None:
62 metadata['document_id'] = document_id
64 return {
65 'content': content,
66 'metadata': metadata,
67 'distance': score,
68 }
71def _source_config(knowledge: KnowledgeModel) -> dict:
72 external = (knowledge.meta or {}).get('external', {})
73 source = external.get('source') or {}
74 return source.get('config') or {}
77def _root_field(path: Optional[str]) -> Optional[str]:
78 if not path:
79 return None
80 return path.split('.')[0]
83def _safe_identifier(value: str, label: str) -> str:
84 if not value or not IDENTIFIER_RE.match(value):
85 raise RuntimeError(f'Invalid {label}')
86 return value
89async def _retrieve_qdrant(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]:
90 try:
91 from qdrant_client import QdrantClient
92 except ImportError as exc:
93 raise RuntimeError('qdrant-client is not installed') from exc
95 if not embedding_function:
96 raise RuntimeError('Embedding function is not configured')
98 config = connection.get('config') or {}
99 external = (knowledge.meta or {}).get('external', {})
100 source = external.get('source') or {}
101 collection_name = source.get('name')
102 if not collection_name:
103 raise RuntimeError('External source collection is not configured')
104 source_config = _source_config(knowledge)
105 vector_field = source_config.get('vector_field') or None
107 vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX)
109 def _search():
110 client = QdrantClient(
111 url=connection.get('endpoint'),
112 api_key=(auth_config or {}).get('api_key'),
113 timeout=config.get('timeout') or 30,
114 )
115 return client.query_points(
116 collection_name=collection_name,
117 query=vector,
118 using=vector_field,
119 limit=count,
120 )
122 response = await asyncio.to_thread(_search)
123 mapping = {
124 'content_field': source_config.get('content_field') or 'payload.text',
125 'metadata_field': source_config.get('metadata_field') or 'payload.metadata',
126 'document_id_field': source_config.get('document_id_field') or 'id',
127 'score_field': 'score',
128 }
130 normalized = []
131 for point in response.points:
132 normalized.append(_normalize_result(point.model_dump(), mapping, knowledge, distance=point.score))
133 return normalized
136async def _retrieve_milvus(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]:
137 try:
138 from pymilvus import MilvusClient
139 except ImportError as exc:
140 raise RuntimeError('pymilvus is not installed') from exc
142 if not embedding_function:
143 raise RuntimeError('Embedding function is not configured')
145 config = connection.get('config') or {}
146 external = (knowledge.meta or {}).get('external', {})
147 source = external.get('source') or {}
148 collection_name = source.get('name')
149 if not collection_name:
150 raise RuntimeError('Milvus collection is not configured')
151 source_config = _source_config(knowledge)
152 vector_field = source_config.get('vector_field') or 'vector'
153 content_field = source_config.get('content_field') or 'data.text'
154 metadata_field = source_config.get('metadata_field') or 'metadata'
156 vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX)
158 def _search():
159 client_kwargs = {
160 'uri': connection.get('endpoint'),
161 }
162 token = (auth_config or {}).get('api_key') or (auth_config or {}).get('token')
163 if token:
164 client_kwargs['token'] = token
165 if config.get('db_name'):
166 client_kwargs['db_name'] = config.get('db_name')
168 client = MilvusClient(**client_kwargs)
169 output_fields = {
170 field
171 for field in (
172 _root_field(content_field),
173 _root_field(metadata_field),
174 _root_field(source_config.get('document_id_field')),
175 )
176 if field and field != vector_field
177 }
178 kwargs = {
179 'collection_name': collection_name,
180 'data': [vector],
181 'anns_field': vector_field,
182 'limit': count,
183 'output_fields': list(output_fields),
184 }
185 return client.search(**kwargs)
187 response = await asyncio.to_thread(_search)
188 mapping = {
189 'content_field': content_field,
190 'metadata_field': metadata_field,
191 'document_id_field': source_config.get('document_id_field') or 'id',
192 'score_field': 'distance',
193 }
195 normalized = []
196 for hit in response[0] if response else []:
197 item = dict(hit)
198 entity = item.get('entity') or {}
199 result = {
200 **entity,
201 'id': item.get('id') or entity.get('id'),
202 'distance': item.get('distance'),
203 }
204 normalized.append(_normalize_result(result, mapping, knowledge, distance=item.get('distance')))
205 return normalized
208async def _retrieve_pgvector(connection, auth_config, knowledge, query, count, embedding_function) -> list[dict]:
209 try:
210 import psycopg
211 from pgvector.psycopg import register_vector
212 from psycopg.rows import dict_row
213 except ImportError as exc:
214 raise RuntimeError('psycopg and pgvector are required for pgvector retrieval') from exc
216 if not embedding_function:
217 raise RuntimeError('Embedding function is not configured')
219 config = connection.get('config') or {}
220 external = (knowledge.meta or {}).get('external', {})
221 source = external.get('source') or {}
222 collection_name = source.get('name')
223 if not collection_name:
224 raise RuntimeError('pgvector collection is not configured')
225 source_config = _source_config(knowledge)
226 table_name = source_config.get('table_name') or 'document_chunk'
227 collection_field = source_config.get('collection_field') or 'collection_name'
228 content_field = source_config.get('content_field') or 'text'
229 vector_field = source_config.get('vector_field') or 'vector'
230 metadata_field = source_config.get('metadata_field') or 'vmetadata'
231 document_id_field = source_config.get('document_id_field') or 'id'
233 vector = await embedding_function(query, prefix=RAG_EMBEDDING_QUERY_PREFIX)
235 def _search():
236 from psycopg import sql
238 table_identifier = sql.SQL('.').join(
239 sql.Identifier(_safe_identifier(part, 'table name')) for part in table_name.split('.')
240 )
241 collection_identifier = sql.Identifier(_safe_identifier(collection_field, 'collection field'))
242 content_identifier = sql.Identifier(_safe_identifier(content_field, 'content field'))
243 vector_identifier = sql.Identifier(_safe_identifier(vector_field, 'vector field'))
244 document_id_identifier = sql.Identifier(_safe_identifier(document_id_field, 'document id field'))
245 metadata_sql = (
246 sql.Identifier(_safe_identifier(metadata_field, 'metadata field'))
247 if metadata_field
248 else sql.SQL("'{}'::jsonb")
249 )
251 with psycopg.connect(
252 connection.get('endpoint'),
253 row_factory=dict_row,
254 connect_timeout=config.get('timeout') or 30,
255 ) as conn:
256 register_vector(conn)
257 with conn.cursor() as cur:
258 cur.execute(
259 sql.SQL(
260 """
261 SELECT {document_id} AS id,
262 {content} AS content,
263 {metadata} AS metadata,
264 {vector_column} <=> %s AS distance
265 FROM {table_name}
266 WHERE {collection} = %s
267 ORDER BY distance ASC
268 LIMIT %s
269 """
270 ).format(
271 document_id=document_id_identifier,
272 content=content_identifier,
273 metadata=metadata_sql,
274 vector_column=vector_identifier,
275 table_name=table_identifier,
276 collection=collection_identifier,
277 ),
278 (vector, collection_name, count),
279 )
280 return cur.fetchall()
282 rows = await asyncio.to_thread(_search)
283 mapping = {
284 'content_field': 'content',
285 'metadata_field': 'metadata',
286 'document_id_field': 'id',
287 'score_field': 'distance',
288 }
289 return [_normalize_result(row, mapping, knowledge, distance=row.get('distance')) for row in rows]
292async def retrieve_external_knowledge(
293 request,
294 knowledge: KnowledgeModel,
295 queries: list[str],
296 count: int,
297 user=None,
298) -> dict:
299 external = (knowledge.meta or {}).get('external', {})
300 connection_id = external.get('connection_id')
301 if not connection_id:
302 raise RuntimeError('External knowledge connection is not configured')
304 connection = await _get_external_connection(connection_id)
305 if not connection:
306 raise RuntimeError('External knowledge connection not found')
308 return await retrieve_external_knowledge_for_connection(request, knowledge, connection, queries, count, user=user)
311async def retrieve_external_knowledge_for_connection(
312 request,
313 knowledge: KnowledgeModel,
314 connection: dict,
315 queries: list[str],
316 count: int,
317 user=None,
318) -> dict:
319 auth_config = connection.get('auth_config') or {}
320 if not connection.get('enabled', True):
321 raise RuntimeError('External knowledge connection is disabled')
323 started_at = time.monotonic()
324 chunks = []
325 provider = (connection.get('provider') or '').lower()
327 for query in queries:
328 if provider == 'qdrant':
329 chunks.extend(
330 await _retrieve_qdrant(
331 connection,
332 auth_config,
333 knowledge,
334 query,
335 count,
336 getattr(request.app.state, 'EMBEDDING_FUNCTION', None),
337 )
338 )
339 elif provider == 'milvus':
340 chunks.extend(
341 await _retrieve_milvus(
342 connection,
343 auth_config,
344 knowledge,
345 query,
346 count,
347 getattr(request.app.state, 'EMBEDDING_FUNCTION', None),
348 )
349 )
350 elif provider == 'pgvector':
351 chunks.extend(
352 await _retrieve_pgvector(
353 connection,
354 auth_config,
355 knowledge,
356 query,
357 count,
358 getattr(request.app.state, 'EMBEDDING_FUNCTION', None),
359 )
360 )
361 else:
362 raise RuntimeError(f'Unsupported external knowledge provider: {connection.get("provider")}')
364 chunks = chunks[:count]
365 log.info(
366 'external_knowledge_retrieval knowledge_id=%s connection_id=%s provider=%s user_id=%s latency_ms=%s result_count=%s',
367 knowledge.id,
368 connection.get('id'),
369 connection.get('provider'),
370 getattr(user, 'id', None),
371 round((time.monotonic() - started_at) * 1000),
372 len(chunks),
373 )
375 return {
376 'documents': [[chunk['content'] for chunk in chunks]],
377 'metadatas': [[chunk['metadata'] for chunk in chunks]],
378 'distances': [[chunk['distance'] for chunk in chunks]],
379 }