chore: format
This commit is contained in:
@@ -850,7 +850,11 @@ def load_oauth_providers():
|
||||
if FEISHU_CLIENT_ID.value:
|
||||
configured_providers.append("Feishu")
|
||||
|
||||
if configured_providers and not OPENID_PROVIDER_URL.value and not OPENID_END_SESSION_ENDPOINT.value:
|
||||
if (
|
||||
configured_providers
|
||||
and not OPENID_PROVIDER_URL.value
|
||||
and not OPENID_END_SESSION_ENDPOINT.value
|
||||
):
|
||||
provider_list = ", ".join(configured_providers)
|
||||
log.warning(
|
||||
f"⚠️ OAuth providers configured ({provider_list}) but OPENID_PROVIDER_URL not set - logout will not work!"
|
||||
@@ -2371,7 +2375,8 @@ if VECTOR_DB == "chroma":
|
||||
MARIADB_VECTOR_DB_URL = os.environ.get("MARIADB_VECTOR_DB_URL", "").strip()
|
||||
|
||||
MARIADB_VECTOR_INITIALIZE_MAX_VECTOR_LENGTH = int(
|
||||
os.environ.get("MARIADB_VECTOR_INITIALIZE_MAX_VECTOR_LENGTH", "1536").strip() or "1536"
|
||||
os.environ.get("MARIADB_VECTOR_INITIALIZE_MAX_VECTOR_LENGTH", "1536").strip()
|
||||
or "1536"
|
||||
)
|
||||
|
||||
# Distance strategy:
|
||||
@@ -2382,7 +2387,9 @@ MARIADB_VECTOR_DISTANCE_STRATEGY = (
|
||||
)
|
||||
|
||||
# HNSW M parameter (MariaDB VECTOR INDEX ... M=<int>)
|
||||
MARIADB_VECTOR_INDEX_M = int(os.environ.get("MARIADB_VECTOR_INDEX_M", "8").strip() or "8")
|
||||
MARIADB_VECTOR_INDEX_M = int(
|
||||
os.environ.get("MARIADB_VECTOR_INDEX_M", "8").strip() or "8"
|
||||
)
|
||||
|
||||
# Pooling (MariaDB-Vector)
|
||||
MARIADB_VECTOR_POOL_SIZE = os.environ.get("MARIADB_VECTOR_POOL_SIZE", None)
|
||||
|
||||
@@ -590,7 +590,9 @@ def generate_openai_batch_embeddings(
|
||||
if "data" in data:
|
||||
return [elem["embedding"] for elem in data["data"]]
|
||||
else:
|
||||
raise ValueError("Unexpected OpenAI embeddings response: missing 'data' key")
|
||||
raise ValueError(
|
||||
"Unexpected OpenAI embeddings response: missing 'data' key"
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception(f"Error generating openai batch embeddings: {e}")
|
||||
return None
|
||||
@@ -767,7 +769,9 @@ def generate_ollama_batch_embeddings(
|
||||
if "embeddings" in data:
|
||||
return data["embeddings"]
|
||||
else:
|
||||
raise ValueError("Unexpected Ollama embeddings response: missing 'embeddings' key")
|
||||
raise ValueError(
|
||||
"Unexpected Ollama embeddings response: missing 'embeddings' key"
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception(f"Error generating ollama batch embeddings: {e}")
|
||||
return None
|
||||
|
||||
@@ -22,7 +22,12 @@ from open_webui.config import (
|
||||
MARIADB_VECTOR_POOL_TIMEOUT,
|
||||
MARIADB_VECTOR_POOL_RECYCLE,
|
||||
)
|
||||
from open_webui.retrieval.vector.main import GetResult, SearchResult, VectorDBBase, VectorItem
|
||||
from open_webui.retrieval.vector.main import (
|
||||
GetResult,
|
||||
SearchResult,
|
||||
VectorDBBase,
|
||||
VectorItem,
|
||||
)
|
||||
from open_webui.retrieval.vector.utils import process_metadata
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -157,8 +162,7 @@ class MariaDBVectorClient(VectorDBBase):
|
||||
with conn.cursor() as cur:
|
||||
try:
|
||||
dist = self.distance_strategy
|
||||
cur.execute(
|
||||
f"""
|
||||
cur.execute(f"""
|
||||
CREATE TABLE IF NOT EXISTS document_chunk (
|
||||
-- MariaDB Vector requires the table PRIMARY KEY used with a VECTOR INDEX to be <= 256 bytes.
|
||||
-- VARCHAR has internal length/metadata overhead, so VARCHAR(255) can exceed the 256-byte limit.
|
||||
@@ -173,8 +177,7 @@ class MariaDBVectorClient(VectorDBBase):
|
||||
VECTOR INDEX (embedding) M={self.index_m} DISTANCE={dist},
|
||||
INDEX idx_document_chunk_collection_name (collection_name)
|
||||
) ENGINE=InnoDB;
|
||||
"""
|
||||
)
|
||||
""")
|
||||
conn.commit()
|
||||
except Exception as e:
|
||||
conn.rollback()
|
||||
@@ -220,7 +223,11 @@ class MariaDBVectorClient(VectorDBBase):
|
||||
"""
|
||||
Return the MariaDB Vector distance function name for the configured strategy.
|
||||
"""
|
||||
return "vec_distance_cosine" if self.distance_strategy == "cosine" else "vec_distance_euclidean"
|
||||
return (
|
||||
"vec_distance_cosine"
|
||||
if self.distance_strategy == "cosine"
|
||||
else "vec_distance_euclidean"
|
||||
)
|
||||
|
||||
def _score_from_dist(self, dist: float) -> float:
|
||||
"""
|
||||
@@ -442,12 +449,19 @@ class MariaDBVectorClient(VectorDBBase):
|
||||
documents[q_idx].append(rtext)
|
||||
metadatas[q_idx].append(_safe_json(rmeta))
|
||||
|
||||
return SearchResult(ids=ids, distances=distances, documents=documents, metadatas=metadatas)
|
||||
return SearchResult(
|
||||
ids=ids,
|
||||
distances=distances,
|
||||
documents=documents,
|
||||
metadatas=metadatas,
|
||||
)
|
||||
except Exception as e:
|
||||
log.exception(f"[MARIADB_VECTOR] search() failed: {e}")
|
||||
return None
|
||||
|
||||
def query(self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None) -> Optional[GetResult]:
|
||||
def query(
|
||||
self, collection_name: str, filter: Dict[str, Any], limit: Optional[int] = None
|
||||
) -> Optional[GetResult]:
|
||||
"""
|
||||
Retrieve documents by metadata filter (non-vector query).
|
||||
"""
|
||||
@@ -472,7 +486,9 @@ class MariaDBVectorClient(VectorDBBase):
|
||||
metadatas = [[_safe_json(r[2]) for r in rows]]
|
||||
return GetResult(ids=ids, documents=documents, metadatas=metadatas)
|
||||
|
||||
def get(self, collection_name: str, limit: Optional[int] = None) -> Optional[GetResult]:
|
||||
def get(
|
||||
self, collection_name: str, limit: Optional[int] = None
|
||||
) -> Optional[GetResult]:
|
||||
"""
|
||||
Retrieve documents in a collection without filtering (optionally limited).
|
||||
"""
|
||||
@@ -549,7 +565,10 @@ class MariaDBVectorClient(VectorDBBase):
|
||||
try:
|
||||
with self._connect() as conn:
|
||||
with conn.cursor() as cur:
|
||||
cur.execute("SELECT 1 FROM document_chunk WHERE collection_name = ? LIMIT 1", (collection_name,))
|
||||
cur.execute(
|
||||
"SELECT 1 FROM document_chunk WHERE collection_name = ? LIMIT 1",
|
||||
(collection_name,),
|
||||
)
|
||||
return cur.fetchone() is not None
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@@ -58,7 +58,9 @@ class Vector:
|
||||
|
||||
return OpenGaussClient()
|
||||
case VectorType.MARIADB_VECTOR:
|
||||
from open_webui.retrieval.vector.dbs.mariadb_vector import MariaDBVectorClient
|
||||
from open_webui.retrieval.vector.dbs.mariadb_vector import (
|
||||
MariaDBVectorClient,
|
||||
)
|
||||
|
||||
return MariaDBVectorClient()
|
||||
case VectorType.ELASTICSEARCH:
|
||||
|
||||
@@ -583,6 +583,7 @@ def update_file_data_content_by_id(
|
||||
request,
|
||||
ProcessFileForm(file_id=id, content=form_data.content),
|
||||
user=user,
|
||||
db=db,
|
||||
)
|
||||
file = Files.get_file_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
|
||||
@@ -1747,7 +1747,9 @@ async def download_file_stream(
|
||||
|
||||
yield f"data: {json.dumps(res)}\n\n"
|
||||
else:
|
||||
raise RuntimeError("Ollama: Could not create blob, Please try again.")
|
||||
raise RuntimeError(
|
||||
"Ollama: Could not create blob, Please try again."
|
||||
)
|
||||
|
||||
|
||||
# url = "https://huggingface.co/TheBloke/stablelm-zephyr-3b-GGUF/resolve/main/stablelm-zephyr-3b.Q2_K.gguf"
|
||||
|
||||
@@ -1706,7 +1706,9 @@ class OAuthManager:
|
||||
redirect_url = f"{redirect_base_url}/auth"
|
||||
|
||||
if error_message:
|
||||
redirect_url = f"{redirect_url}?error={urllib.parse.quote_plus(error_message)}"
|
||||
redirect_url = (
|
||||
f"{redirect_url}?error={urllib.parse.quote_plus(error_message)}"
|
||||
)
|
||||
return RedirectResponse(url=redirect_url, headers=response.headers)
|
||||
|
||||
response = RedirectResponse(url=redirect_url, headers=response.headers)
|
||||
|
||||
@@ -162,9 +162,7 @@ def truncate_content(content: str, max_chars: int, mode: str = "middletruncate")
|
||||
return f"{content[:half]}...{content[-(max_chars - half):]}"
|
||||
|
||||
|
||||
def apply_content_filter(
|
||||
messages: list[dict], filter_str: str
|
||||
) -> list[dict]:
|
||||
def apply_content_filter(messages: list[dict], filter_str: str) -> list[dict]:
|
||||
"""Apply a content filter to each message's content.
|
||||
|
||||
filter_str is like 'middletruncate:500', 'start:200', or 'end:200'.
|
||||
@@ -238,9 +236,7 @@ def replace_messages_variable(
|
||||
else:
|
||||
half = mid // 2
|
||||
start_msgs = messages[:half]
|
||||
end_msgs = (
|
||||
messages[-half:] if mid % 2 == 0 else messages[-(half + 1) :]
|
||||
)
|
||||
end_msgs = messages[-half:] if mid % 2 == 0 else messages[-(half + 1) :]
|
||||
selected = start_msgs + end_msgs
|
||||
content_filter = middle_filter
|
||||
else:
|
||||
|
||||
@@ -32,8 +32,7 @@ async def post_webhook(name: str, url: str, message: str, event_data: dict) -> b
|
||||
else:
|
||||
user_dict = json.loads(user_data)
|
||||
facts = [
|
||||
{"name": name, "value": value}
|
||||
for name, value in user_dict.items()
|
||||
{"name": name, "value": value} for name, value in user_dict.items()
|
||||
]
|
||||
payload = {
|
||||
"@type": "MessageCard",
|
||||
|
||||
Reference in New Issue
Block a user