mirror of
https://github.com/langgenius/dify.git
synced 2026-08-28 21:23:22 +08:00
Merge branch 'main' into feat/mcp-06-18
This commit is contained in:
commit
a538f80e95
2
.github/workflows/autofix.yml
vendored
2
.github/workflows/autofix.yml
vendored
@ -30,6 +30,8 @@ jobs:
|
|||||||
run: |
|
run: |
|
||||||
uvx --from ast-grep-cli sg --pattern 'db.session.query($WHATEVER).filter($HERE)' --rewrite 'db.session.query($WHATEVER).where($HERE)' -l py --update-all
|
uvx --from ast-grep-cli sg --pattern 'db.session.query($WHATEVER).filter($HERE)' --rewrite 'db.session.query($WHATEVER).where($HERE)' -l py --update-all
|
||||||
uvx --from ast-grep-cli sg --pattern 'session.query($WHATEVER).filter($HERE)' --rewrite 'session.query($WHATEVER).where($HERE)' -l py --update-all
|
uvx --from ast-grep-cli sg --pattern 'session.query($WHATEVER).filter($HERE)' --rewrite 'session.query($WHATEVER).where($HERE)' -l py --update-all
|
||||||
|
uvx --from ast-grep-cli sg -p '$A = db.Column($$$B)' -r '$A = mapped_column($$$B)' -l py --update-all
|
||||||
|
uvx --from ast-grep-cli sg -p '$A : $T = db.Column($$$B)' -r '$A : $T = mapped_column($$$B)' -l py --update-all
|
||||||
# Convert Optional[T] to T | None (ignoring quoted types)
|
# Convert Optional[T] to T | None (ignoring quoted types)
|
||||||
cat > /tmp/optional-rule.yml << 'EOF'
|
cat > /tmp/optional-rule.yml << 'EOF'
|
||||||
id: convert-optional-to-union
|
id: convert-optional-to-union
|
||||||
|
|||||||
3
.github/workflows/build-push.yml
vendored
3
.github/workflows/build-push.yml
vendored
@ -4,8 +4,7 @@ on:
|
|||||||
push:
|
push:
|
||||||
branches:
|
branches:
|
||||||
- "main"
|
- "main"
|
||||||
- "deploy/dev"
|
- "deploy/**"
|
||||||
- "deploy/enterprise"
|
|
||||||
- "build/**"
|
- "build/**"
|
||||||
- "release/e-*"
|
- "release/e-*"
|
||||||
- "hotfix/**"
|
- "hotfix/**"
|
||||||
|
|||||||
2
.github/workflows/deploy-dev.yml
vendored
2
.github/workflows/deploy-dev.yml
vendored
@ -18,7 +18,7 @@ jobs:
|
|||||||
- name: Deploy to server
|
- name: Deploy to server
|
||||||
uses: appleboy/ssh-action@v0.1.8
|
uses: appleboy/ssh-action@v0.1.8
|
||||||
with:
|
with:
|
||||||
host: ${{ secrets.RAG_SSH_HOST }}
|
host: ${{ secrets.SSH_HOST }}
|
||||||
username: ${{ secrets.SSH_USER }}
|
username: ${{ secrets.SSH_USER }}
|
||||||
key: ${{ secrets.SSH_PRIVATE_KEY }}
|
key: ${{ secrets.SSH_PRIVATE_KEY }}
|
||||||
script: |
|
script: |
|
||||||
|
|||||||
@ -1,4 +1,4 @@
|
|||||||
name: Deploy RAG Dev
|
name: Deploy Trigger Dev
|
||||||
|
|
||||||
permissions:
|
permissions:
|
||||||
contents: read
|
contents: read
|
||||||
@ -7,7 +7,7 @@ on:
|
|||||||
workflow_run:
|
workflow_run:
|
||||||
workflows: ["Build and Push API & Web"]
|
workflows: ["Build and Push API & Web"]
|
||||||
branches:
|
branches:
|
||||||
- "deploy/rag-dev"
|
- "deploy/trigger-dev"
|
||||||
types:
|
types:
|
||||||
- completed
|
- completed
|
||||||
|
|
||||||
@ -16,12 +16,12 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: |
|
if: |
|
||||||
github.event.workflow_run.conclusion == 'success' &&
|
github.event.workflow_run.conclusion == 'success' &&
|
||||||
github.event.workflow_run.head_branch == 'deploy/rag-dev'
|
github.event.workflow_run.head_branch == 'deploy/trigger-dev'
|
||||||
steps:
|
steps:
|
||||||
- name: Deploy to server
|
- name: Deploy to server
|
||||||
uses: appleboy/ssh-action@v0.1.8
|
uses: appleboy/ssh-action@v0.1.8
|
||||||
with:
|
with:
|
||||||
host: ${{ secrets.RAG_SSH_HOST }}
|
host: ${{ secrets.TRIGGER_SSH_HOST }}
|
||||||
username: ${{ secrets.SSH_USER }}
|
username: ${{ secrets.SSH_USER }}
|
||||||
key: ${{ secrets.SSH_PRIVATE_KEY }}
|
key: ${{ secrets.SSH_PRIVATE_KEY }}
|
||||||
script: |
|
script: |
|
||||||
@ -343,6 +343,15 @@ OCEANBASE_VECTOR_DATABASE=test
|
|||||||
OCEANBASE_MEMORY_LIMIT=6G
|
OCEANBASE_MEMORY_LIMIT=6G
|
||||||
OCEANBASE_ENABLE_HYBRID_SEARCH=false
|
OCEANBASE_ENABLE_HYBRID_SEARCH=false
|
||||||
|
|
||||||
|
# AlibabaCloud MySQL Vector configuration
|
||||||
|
ALIBABACLOUD_MYSQL_HOST=127.0.0.1
|
||||||
|
ALIBABACLOUD_MYSQL_PORT=3306
|
||||||
|
ALIBABACLOUD_MYSQL_USER=root
|
||||||
|
ALIBABACLOUD_MYSQL_PASSWORD=root
|
||||||
|
ALIBABACLOUD_MYSQL_DATABASE=dify
|
||||||
|
ALIBABACLOUD_MYSQL_MAX_CONNECTION=5
|
||||||
|
ALIBABACLOUD_MYSQL_HNSW_M=6
|
||||||
|
|
||||||
# openGauss configuration
|
# openGauss configuration
|
||||||
OPENGAUSS_HOST=127.0.0.1
|
OPENGAUSS_HOST=127.0.0.1
|
||||||
OPENGAUSS_PORT=6600
|
OPENGAUSS_PORT=6600
|
||||||
|
|||||||
@ -81,7 +81,6 @@ ignore = [
|
|||||||
"SIM113", # enumerate-for-loop
|
"SIM113", # enumerate-for-loop
|
||||||
"SIM117", # multiple-with-statements
|
"SIM117", # multiple-with-statements
|
||||||
"SIM210", # if-expr-with-true-false
|
"SIM210", # if-expr-with-true-false
|
||||||
"UP038", # deprecated and not recommended by Ruff, https://docs.astral.sh/ruff/rules/non-pep604-isinstance/
|
|
||||||
]
|
]
|
||||||
|
|
||||||
[lint.per-file-ignores]
|
[lint.per-file-ignores]
|
||||||
|
|||||||
@ -1521,6 +1521,14 @@ def transform_datasource_credentials():
|
|||||||
auth_count = 0
|
auth_count = 0
|
||||||
for firecrawl_tenant_credential in firecrawl_tenant_credentials:
|
for firecrawl_tenant_credential in firecrawl_tenant_credentials:
|
||||||
auth_count += 1
|
auth_count += 1
|
||||||
|
if not firecrawl_tenant_credential.credentials:
|
||||||
|
click.echo(
|
||||||
|
click.style(
|
||||||
|
f"Skipping firecrawl credential for tenant {tenant_id} due to missing credentials.",
|
||||||
|
fg="yellow",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
# get credential api key
|
# get credential api key
|
||||||
credentials_json = json.loads(firecrawl_tenant_credential.credentials)
|
credentials_json = json.loads(firecrawl_tenant_credential.credentials)
|
||||||
api_key = credentials_json.get("config", {}).get("api_key")
|
api_key = credentials_json.get("config", {}).get("api_key")
|
||||||
@ -1576,6 +1584,14 @@ def transform_datasource_credentials():
|
|||||||
auth_count = 0
|
auth_count = 0
|
||||||
for jina_tenant_credential in jina_tenant_credentials:
|
for jina_tenant_credential in jina_tenant_credentials:
|
||||||
auth_count += 1
|
auth_count += 1
|
||||||
|
if not jina_tenant_credential.credentials:
|
||||||
|
click.echo(
|
||||||
|
click.style(
|
||||||
|
f"Skipping jina credential for tenant {tenant_id} due to missing credentials.",
|
||||||
|
fg="yellow",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
continue
|
||||||
# get credential api key
|
# get credential api key
|
||||||
credentials_json = json.loads(jina_tenant_credential.credentials)
|
credentials_json = json.loads(jina_tenant_credential.credentials)
|
||||||
api_key = credentials_json.get("config", {}).get("api_key")
|
api_key = credentials_json.get("config", {}).get("api_key")
|
||||||
|
|||||||
@ -18,6 +18,7 @@ from .storage.opendal_storage_config import OpenDALStorageConfig
|
|||||||
from .storage.supabase_storage_config import SupabaseStorageConfig
|
from .storage.supabase_storage_config import SupabaseStorageConfig
|
||||||
from .storage.tencent_cos_storage_config import TencentCloudCOSStorageConfig
|
from .storage.tencent_cos_storage_config import TencentCloudCOSStorageConfig
|
||||||
from .storage.volcengine_tos_storage_config import VolcengineTOSStorageConfig
|
from .storage.volcengine_tos_storage_config import VolcengineTOSStorageConfig
|
||||||
|
from .vdb.alibabacloud_mysql_config import AlibabaCloudMySQLConfig
|
||||||
from .vdb.analyticdb_config import AnalyticdbConfig
|
from .vdb.analyticdb_config import AnalyticdbConfig
|
||||||
from .vdb.baidu_vector_config import BaiduVectorDBConfig
|
from .vdb.baidu_vector_config import BaiduVectorDBConfig
|
||||||
from .vdb.chroma_config import ChromaConfig
|
from .vdb.chroma_config import ChromaConfig
|
||||||
@ -330,6 +331,7 @@ class MiddlewareConfig(
|
|||||||
ClickzettaConfig,
|
ClickzettaConfig,
|
||||||
HuaweiCloudConfig,
|
HuaweiCloudConfig,
|
||||||
MilvusConfig,
|
MilvusConfig,
|
||||||
|
AlibabaCloudMySQLConfig,
|
||||||
MyScaleConfig,
|
MyScaleConfig,
|
||||||
OpenSearchConfig,
|
OpenSearchConfig,
|
||||||
OracleConfig,
|
OracleConfig,
|
||||||
|
|||||||
54
api/configs/middleware/vdb/alibabacloud_mysql_config.py
Normal file
54
api/configs/middleware/vdb/alibabacloud_mysql_config.py
Normal file
@ -0,0 +1,54 @@
|
|||||||
|
from pydantic import Field, PositiveInt
|
||||||
|
from pydantic_settings import BaseSettings
|
||||||
|
|
||||||
|
|
||||||
|
class AlibabaCloudMySQLConfig(BaseSettings):
|
||||||
|
"""
|
||||||
|
Configuration settings for AlibabaCloud MySQL vector database
|
||||||
|
"""
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_HOST: str = Field(
|
||||||
|
description="Hostname or IP address of the AlibabaCloud MySQL server (e.g., 'localhost' or 'mysql.aliyun.com')",
|
||||||
|
default="localhost",
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_PORT: PositiveInt = Field(
|
||||||
|
description="Port number on which the AlibabaCloud MySQL server is listening (default is 3306)",
|
||||||
|
default=3306,
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_USER: str = Field(
|
||||||
|
description="Username for authenticating with AlibabaCloud MySQL (default is 'root')",
|
||||||
|
default="root",
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_PASSWORD: str = Field(
|
||||||
|
description="Password for authenticating with AlibabaCloud MySQL (default is an empty string)",
|
||||||
|
default="",
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_DATABASE: str = Field(
|
||||||
|
description="Name of the AlibabaCloud MySQL database to connect to (default is 'dify')",
|
||||||
|
default="dify",
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_MAX_CONNECTION: PositiveInt = Field(
|
||||||
|
description="Maximum number of connections in the connection pool",
|
||||||
|
default=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_CHARSET: str = Field(
|
||||||
|
description="Character set for AlibabaCloud MySQL connection (default is 'utf8mb4')",
|
||||||
|
default="utf8mb4",
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_DISTANCE_FUNCTION: str = Field(
|
||||||
|
description="Distance function used for vector similarity search in AlibabaCloud MySQL "
|
||||||
|
"(e.g., 'cosine', 'euclidean')",
|
||||||
|
default="cosine",
|
||||||
|
)
|
||||||
|
|
||||||
|
ALIBABACLOUD_MYSQL_HNSW_M: PositiveInt = Field(
|
||||||
|
description="Maximum number of connections per layer for HNSW vector index (default is 6, range: 3-200)",
|
||||||
|
default=6,
|
||||||
|
)
|
||||||
@ -1,23 +1,24 @@
|
|||||||
from enum import Enum
|
from enum import StrEnum
|
||||||
from typing import Literal
|
from typing import Literal
|
||||||
|
|
||||||
from pydantic import Field, PositiveInt
|
from pydantic import Field, PositiveInt
|
||||||
from pydantic_settings import BaseSettings
|
from pydantic_settings import BaseSettings
|
||||||
|
|
||||||
|
|
||||||
|
class AuthMethod(StrEnum):
|
||||||
|
"""
|
||||||
|
Authentication method for OpenSearch
|
||||||
|
"""
|
||||||
|
|
||||||
|
BASIC = "basic"
|
||||||
|
AWS_MANAGED_IAM = "aws_managed_iam"
|
||||||
|
|
||||||
|
|
||||||
class OpenSearchConfig(BaseSettings):
|
class OpenSearchConfig(BaseSettings):
|
||||||
"""
|
"""
|
||||||
Configuration settings for OpenSearch
|
Configuration settings for OpenSearch
|
||||||
"""
|
"""
|
||||||
|
|
||||||
class AuthMethod(Enum):
|
|
||||||
"""
|
|
||||||
Authentication method for OpenSearch
|
|
||||||
"""
|
|
||||||
|
|
||||||
BASIC = "basic"
|
|
||||||
AWS_MANAGED_IAM = "aws_managed_iam"
|
|
||||||
|
|
||||||
OPENSEARCH_HOST: str | None = Field(
|
OPENSEARCH_HOST: str | None = Field(
|
||||||
description="Hostname or IP address of the OpenSearch server (e.g., 'localhost' or 'opensearch.example.com')",
|
description="Hostname or IP address of the OpenSearch server (e.g., 'localhost' or 'opensearch.example.com')",
|
||||||
default=None,
|
default=None,
|
||||||
|
|||||||
@ -1,5 +1,4 @@
|
|||||||
import flask_restx
|
import flask_restx
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, fields, marshal_with
|
from flask_restx import Resource, fields, marshal_with
|
||||||
from flask_restx._http import HTTPStatus
|
from flask_restx._http import HTTPStatus
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
@ -8,7 +7,8 @@ from werkzeug.exceptions import Forbidden
|
|||||||
|
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from libs.helper import TimestampField
|
from libs.helper import TimestampField
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from models.dataset import Dataset
|
from models.dataset import Dataset
|
||||||
from models.model import ApiToken, App
|
from models.model import ApiToken, App
|
||||||
|
|
||||||
@ -57,6 +57,8 @@ class BaseApiKeyListResource(Resource):
|
|||||||
def get(self, resource_id):
|
def get(self, resource_id):
|
||||||
assert self.resource_id_field is not None, "resource_id_field must be set"
|
assert self.resource_id_field is not None, "resource_id_field must be set"
|
||||||
resource_id = str(resource_id)
|
resource_id = str(resource_id)
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
_get_resource(resource_id, current_user.current_tenant_id, self.resource_model)
|
_get_resource(resource_id, current_user.current_tenant_id, self.resource_model)
|
||||||
keys = db.session.scalars(
|
keys = db.session.scalars(
|
||||||
select(ApiToken).where(
|
select(ApiToken).where(
|
||||||
@ -69,8 +71,10 @@ class BaseApiKeyListResource(Resource):
|
|||||||
def post(self, resource_id):
|
def post(self, resource_id):
|
||||||
assert self.resource_id_field is not None, "resource_id_field must be set"
|
assert self.resource_id_field is not None, "resource_id_field must be set"
|
||||||
resource_id = str(resource_id)
|
resource_id = str(resource_id)
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
_get_resource(resource_id, current_user.current_tenant_id, self.resource_model)
|
_get_resource(resource_id, current_user.current_tenant_id, self.resource_model)
|
||||||
if not current_user.is_editor:
|
if not current_user.has_edit_permission:
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
current_key_count = (
|
current_key_count = (
|
||||||
@ -108,6 +112,8 @@ class BaseApiKeyResource(Resource):
|
|||||||
assert self.resource_id_field is not None, "resource_id_field must be set"
|
assert self.resource_id_field is not None, "resource_id_field must be set"
|
||||||
resource_id = str(resource_id)
|
resource_id = str(resource_id)
|
||||||
api_key_id = str(api_key_id)
|
api_key_id = str(api_key_id)
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
_get_resource(resource_id, current_user.current_tenant_id, self.resource_model)
|
_get_resource(resource_id, current_user.current_tenant_id, self.resource_model)
|
||||||
|
|
||||||
# The role of the current user in the ta table must be admin or owner
|
# The role of the current user in the ta table must be admin or owner
|
||||||
|
|||||||
@ -304,7 +304,7 @@ class AppCopyApi(Resource):
|
|||||||
account = cast(Account, current_user)
|
account = cast(Account, current_user)
|
||||||
result = import_service.import_app(
|
result = import_service.import_app(
|
||||||
account=account,
|
account=account,
|
||||||
import_mode=ImportMode.YAML_CONTENT.value,
|
import_mode=ImportMode.YAML_CONTENT,
|
||||||
yaml_content=yaml_content,
|
yaml_content=yaml_content,
|
||||||
name=args.get("name"),
|
name=args.get("name"),
|
||||||
description=args.get("description"),
|
description=args.get("description"),
|
||||||
|
|||||||
@ -70,9 +70,9 @@ class AppImportApi(Resource):
|
|||||||
EnterpriseService.WebAppAuth.update_app_access_mode(result.app_id, "private")
|
EnterpriseService.WebAppAuth.update_app_access_mode(result.app_id, "private")
|
||||||
# Return appropriate status code based on result
|
# Return appropriate status code based on result
|
||||||
status = result.status
|
status = result.status
|
||||||
if status == ImportStatus.FAILED.value:
|
if status == ImportStatus.FAILED:
|
||||||
return result.model_dump(mode="json"), 400
|
return result.model_dump(mode="json"), 400
|
||||||
elif status == ImportStatus.PENDING.value:
|
elif status == ImportStatus.PENDING:
|
||||||
return result.model_dump(mode="json"), 202
|
return result.model_dump(mode="json"), 202
|
||||||
return result.model_dump(mode="json"), 200
|
return result.model_dump(mode="json"), 200
|
||||||
|
|
||||||
@ -97,7 +97,7 @@ class AppImportConfirmApi(Resource):
|
|||||||
session.commit()
|
session.commit()
|
||||||
|
|
||||||
# Return appropriate status code based on result
|
# Return appropriate status code based on result
|
||||||
if result.status == ImportStatus.FAILED.value:
|
if result.status == ImportStatus.FAILED:
|
||||||
return result.model_dump(mode="json"), 400
|
return result.model_dump(mode="json"), 400
|
||||||
return result.model_dump(mode="json"), 200
|
return result.model_dump(mode="json"), 200
|
||||||
|
|
||||||
|
|||||||
@ -309,7 +309,7 @@ class ChatConversationApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if app_model.mode == AppMode.ADVANCED_CHAT:
|
if app_model.mode == AppMode.ADVANCED_CHAT:
|
||||||
query = query.where(Conversation.invoke_from != InvokeFrom.DEBUGGER.value)
|
query = query.where(Conversation.invoke_from != InvokeFrom.DEBUGGER)
|
||||||
|
|
||||||
match args["sort_by"]:
|
match args["sort_by"]:
|
||||||
case "created_at":
|
case "created_at":
|
||||||
|
|||||||
@ -14,6 +14,7 @@ from core.tools.tool_manager import ToolManager
|
|||||||
from core.tools.utils.configuration import ToolParameterConfigurationManager
|
from core.tools.utils.configuration import ToolParameterConfigurationManager
|
||||||
from events.app_event import app_model_config_was_updated
|
from events.app_event import app_model_config_was_updated
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
|
from libs.datetime_utils import naive_utc_now
|
||||||
from libs.login import login_required
|
from libs.login import login_required
|
||||||
from models.account import Account
|
from models.account import Account
|
||||||
from models.model import AppMode, AppModelConfig
|
from models.model import AppMode, AppModelConfig
|
||||||
@ -90,7 +91,7 @@ class ModelConfigResource(Resource):
|
|||||||
if not isinstance(tool, dict) or len(tool.keys()) <= 3:
|
if not isinstance(tool, dict) or len(tool.keys()) <= 3:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
agent_tool_entity = AgentToolEntity(**tool)
|
agent_tool_entity = AgentToolEntity.model_validate(tool)
|
||||||
# get tool
|
# get tool
|
||||||
try:
|
try:
|
||||||
tool_runtime = ToolManager.get_agent_tool_runtime(
|
tool_runtime = ToolManager.get_agent_tool_runtime(
|
||||||
@ -124,7 +125,7 @@ class ModelConfigResource(Resource):
|
|||||||
# encrypt agent tool parameters if it's secret-input
|
# encrypt agent tool parameters if it's secret-input
|
||||||
agent_mode = new_app_model_config.agent_mode_dict
|
agent_mode = new_app_model_config.agent_mode_dict
|
||||||
for tool in agent_mode.get("tools") or []:
|
for tool in agent_mode.get("tools") or []:
|
||||||
agent_tool_entity = AgentToolEntity(**tool)
|
agent_tool_entity = AgentToolEntity.model_validate(tool)
|
||||||
|
|
||||||
# get tool
|
# get tool
|
||||||
key = f"{agent_tool_entity.provider_id}.{agent_tool_entity.provider_type}.{agent_tool_entity.tool_name}"
|
key = f"{agent_tool_entity.provider_id}.{agent_tool_entity.provider_type}.{agent_tool_entity.tool_name}"
|
||||||
@ -172,6 +173,8 @@ class ModelConfigResource(Resource):
|
|||||||
db.session.flush()
|
db.session.flush()
|
||||||
|
|
||||||
app_model.app_model_config_id = new_app_model_config.id
|
app_model.app_model_config_id = new_app_model_config.id
|
||||||
|
app_model.updated_by = current_user.id
|
||||||
|
app_model.updated_at = naive_utc_now()
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
app_model_config_was_updated.send(app_model, app_model_config=new_app_model_config)
|
app_model_config_was_updated.send(app_model, app_model_config=new_app_model_config)
|
||||||
|
|||||||
@ -52,7 +52,7 @@ FROM
|
|||||||
WHERE
|
WHERE
|
||||||
app_id = :app_id
|
app_id = :app_id
|
||||||
AND invoke_from != :invoke_from"""
|
AND invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
@ -127,7 +127,7 @@ class DailyConversationStatistic(Resource):
|
|||||||
sa.func.count(sa.distinct(Message.conversation_id)).label("conversation_count"),
|
sa.func.count(sa.distinct(Message.conversation_id)).label("conversation_count"),
|
||||||
)
|
)
|
||||||
.select_from(Message)
|
.select_from(Message)
|
||||||
.where(Message.app_id == app_model.id, Message.invoke_from != InvokeFrom.DEBUGGER.value)
|
.where(Message.app_id == app_model.id, Message.invoke_from != InvokeFrom.DEBUGGER)
|
||||||
)
|
)
|
||||||
|
|
||||||
if args["start"]:
|
if args["start"]:
|
||||||
@ -190,7 +190,7 @@ FROM
|
|||||||
WHERE
|
WHERE
|
||||||
app_id = :app_id
|
app_id = :app_id
|
||||||
AND invoke_from != :invoke_from"""
|
AND invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
@ -263,7 +263,7 @@ FROM
|
|||||||
WHERE
|
WHERE
|
||||||
app_id = :app_id
|
app_id = :app_id
|
||||||
AND invoke_from != :invoke_from"""
|
AND invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
@ -345,7 +345,7 @@ FROM
|
|||||||
WHERE
|
WHERE
|
||||||
c.app_id = :app_id
|
c.app_id = :app_id
|
||||||
AND m.invoke_from != :invoke_from"""
|
AND m.invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
@ -432,7 +432,7 @@ LEFT JOIN
|
|||||||
WHERE
|
WHERE
|
||||||
m.app_id = :app_id
|
m.app_id = :app_id
|
||||||
AND m.invoke_from != :invoke_from"""
|
AND m.invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
@ -509,7 +509,7 @@ FROM
|
|||||||
WHERE
|
WHERE
|
||||||
app_id = :app_id
|
app_id = :app_id
|
||||||
AND invoke_from != :invoke_from"""
|
AND invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
@ -584,7 +584,7 @@ FROM
|
|||||||
WHERE
|
WHERE
|
||||||
app_id = :app_id
|
app_id = :app_id
|
||||||
AND invoke_from != :invoke_from"""
|
AND invoke_from != :invoke_from"""
|
||||||
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER.value}
|
arg_dict = {"tz": account.timezone, "app_id": app_model.id, "invoke_from": InvokeFrom.DEBUGGER}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
utc_timezone = pytz.utc
|
utc_timezone = pytz.utc
|
||||||
|
|||||||
@ -25,6 +25,7 @@ from factories import file_factory, variable_factory
|
|||||||
from fields.workflow_fields import workflow_fields, workflow_pagination_fields
|
from fields.workflow_fields import workflow_fields, workflow_pagination_fields
|
||||||
from fields.workflow_run_fields import workflow_run_node_execution_fields
|
from fields.workflow_run_fields import workflow_run_node_execution_fields
|
||||||
from libs import helper
|
from libs import helper
|
||||||
|
from libs.datetime_utils import naive_utc_now
|
||||||
from libs.helper import TimestampField, uuid_value
|
from libs.helper import TimestampField, uuid_value
|
||||||
from libs.login import current_user, login_required
|
from libs.login import current_user, login_required
|
||||||
from models import App
|
from models import App
|
||||||
@ -674,8 +675,12 @@ class PublishedWorkflowApi(Resource):
|
|||||||
marked_comment=args.marked_comment or "",
|
marked_comment=args.marked_comment or "",
|
||||||
)
|
)
|
||||||
|
|
||||||
app_model.workflow_id = workflow.id
|
# Update app_model within the same session to ensure atomicity
|
||||||
db.session.commit() # NOTE: this is necessary for update app_model.workflow_id
|
app_model_in_session = session.get(App, app_model.id)
|
||||||
|
if app_model_in_session:
|
||||||
|
app_model_in_session.workflow_id = workflow.id
|
||||||
|
app_model_in_session.updated_by = current_user.id
|
||||||
|
app_model_in_session.updated_at = naive_utc_now()
|
||||||
|
|
||||||
workflow_created_at = TimestampField().format(workflow.created_at)
|
workflow_created_at = TimestampField().format(workflow.created_at)
|
||||||
|
|
||||||
|
|||||||
@ -47,7 +47,7 @@ WHERE
|
|||||||
arg_dict = {
|
arg_dict = {
|
||||||
"tz": account.timezone,
|
"tz": account.timezone,
|
||||||
"app_id": app_model.id,
|
"app_id": app_model.id,
|
||||||
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN.value,
|
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN,
|
||||||
}
|
}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
@ -115,7 +115,7 @@ WHERE
|
|||||||
arg_dict = {
|
arg_dict = {
|
||||||
"tz": account.timezone,
|
"tz": account.timezone,
|
||||||
"app_id": app_model.id,
|
"app_id": app_model.id,
|
||||||
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN.value,
|
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN,
|
||||||
}
|
}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
@ -183,7 +183,7 @@ WHERE
|
|||||||
arg_dict = {
|
arg_dict = {
|
||||||
"tz": account.timezone,
|
"tz": account.timezone,
|
||||||
"app_id": app_model.id,
|
"app_id": app_model.id,
|
||||||
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN.value,
|
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN,
|
||||||
}
|
}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
@ -269,7 +269,7 @@ GROUP BY
|
|||||||
arg_dict = {
|
arg_dict = {
|
||||||
"tz": account.timezone,
|
"tz": account.timezone,
|
||||||
"app_id": app_model.id,
|
"app_id": app_model.id,
|
||||||
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN.value,
|
"triggered_from": WorkflowRunTriggeredFrom.APP_RUN,
|
||||||
}
|
}
|
||||||
|
|
||||||
timezone = pytz.timezone(account.timezone)
|
timezone = pytz.timezone(account.timezone)
|
||||||
|
|||||||
@ -103,7 +103,7 @@ class ActivateApi(Resource):
|
|||||||
account.interface_language = args["interface_language"]
|
account.interface_language = args["interface_language"]
|
||||||
account.timezone = args["timezone"]
|
account.timezone = args["timezone"]
|
||||||
account.interface_theme = "light"
|
account.interface_theme = "light"
|
||||||
account.status = AccountStatus.ACTIVE.value
|
account.status = AccountStatus.ACTIVE
|
||||||
account.initialized_at = naive_utc_now()
|
account.initialized_at = naive_utc_now()
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
|
|||||||
@ -130,11 +130,11 @@ class OAuthCallback(Resource):
|
|||||||
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message={e.description}")
|
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message={e.description}")
|
||||||
|
|
||||||
# Check account status
|
# Check account status
|
||||||
if account.status == AccountStatus.BANNED.value:
|
if account.status == AccountStatus.BANNED:
|
||||||
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Account is banned.")
|
return redirect(f"{dify_config.CONSOLE_WEB_URL}/signin?message=Account is banned.")
|
||||||
|
|
||||||
if account.status == AccountStatus.PENDING.value:
|
if account.status == AccountStatus.PENDING:
|
||||||
account.status = AccountStatus.ACTIVE.value
|
account.status = AccountStatus.ACTIVE
|
||||||
account.initialized_at = naive_utc_now()
|
account.initialized_at = naive_utc_now()
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
|
|||||||
@ -1,9 +1,9 @@
|
|||||||
from flask import request
|
from flask import request
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, reqparse
|
from flask_restx import Resource, reqparse
|
||||||
|
|
||||||
from libs.helper import extract_remote_ip
|
from libs.helper import extract_remote_ip
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from services.billing_service import BillingService
|
from services.billing_service import BillingService
|
||||||
|
|
||||||
from .. import console_ns
|
from .. import console_ns
|
||||||
@ -17,6 +17,8 @@ class ComplianceApi(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@only_edition_cloud
|
@only_edition_cloud
|
||||||
def get(self):
|
def get(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("doc_name", type=str, required=True, location="args")
|
parser.add_argument("doc_name", type=str, required=True, location="args")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|||||||
@ -15,7 +15,7 @@ from core.datasource.entities.datasource_entities import DatasourceProviderType,
|
|||||||
from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin
|
from core.datasource.online_document.online_document_plugin import OnlineDocumentDatasourcePlugin
|
||||||
from core.indexing_runner import IndexingRunner
|
from core.indexing_runner import IndexingRunner
|
||||||
from core.rag.extractor.entity.datasource_type import DatasourceType
|
from core.rag.extractor.entity.datasource_type import DatasourceType
|
||||||
from core.rag.extractor.entity.extract_setting import ExtractSetting
|
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo
|
||||||
from core.rag.extractor.notion_extractor import NotionExtractor
|
from core.rag.extractor.notion_extractor import NotionExtractor
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from fields.data_source_fields import integrate_list_fields, integrate_notion_info_list_fields
|
from fields.data_source_fields import integrate_list_fields, integrate_notion_info_list_fields
|
||||||
@ -256,14 +256,16 @@ class DataSourceNotionApi(Resource):
|
|||||||
credential_id = notion_info.get("credential_id")
|
credential_id = notion_info.get("credential_id")
|
||||||
for page in notion_info["pages"]:
|
for page in notion_info["pages"]:
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.NOTION.value,
|
datasource_type=DatasourceType.NOTION,
|
||||||
notion_info={
|
notion_info=NotionInfo.model_validate(
|
||||||
"credential_id": credential_id,
|
{
|
||||||
"notion_workspace_id": workspace_id,
|
"credential_id": credential_id,
|
||||||
"notion_obj_id": page["page_id"],
|
"notion_workspace_id": workspace_id,
|
||||||
"notion_page_type": page["type"],
|
"notion_obj_id": page["page_id"],
|
||||||
"tenant_id": current_user.current_tenant_id,
|
"notion_page_type": page["type"],
|
||||||
},
|
"tenant_id": current_user.current_tenant_id,
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=args["doc_form"],
|
document_model=args["doc_form"],
|
||||||
)
|
)
|
||||||
extract_settings.append(extract_setting)
|
extract_settings.append(extract_setting)
|
||||||
|
|||||||
@ -24,7 +24,7 @@ from core.model_runtime.entities.model_entities import ModelType
|
|||||||
from core.provider_manager import ProviderManager
|
from core.provider_manager import ProviderManager
|
||||||
from core.rag.datasource.vdb.vector_type import VectorType
|
from core.rag.datasource.vdb.vector_type import VectorType
|
||||||
from core.rag.extractor.entity.datasource_type import DatasourceType
|
from core.rag.extractor.entity.datasource_type import DatasourceType
|
||||||
from core.rag.extractor.entity.extract_setting import ExtractSetting
|
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo
|
||||||
from core.rag.retrieval.retrieval_methods import RetrievalMethod
|
from core.rag.retrieval.retrieval_methods import RetrievalMethod
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from fields.app_fields import related_app_list
|
from fields.app_fields import related_app_list
|
||||||
@ -45,6 +45,79 @@ def _validate_name(name: str) -> str:
|
|||||||
return name
|
return name
|
||||||
|
|
||||||
|
|
||||||
|
def _get_retrieval_methods_by_vector_type(vector_type: str | None, is_mock: bool = False) -> dict[str, list[str]]:
|
||||||
|
"""
|
||||||
|
Get supported retrieval methods based on vector database type.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
vector_type: Vector database type, can be None
|
||||||
|
is_mock: Whether this is a Mock API, affects MILVUS handling
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary containing supported retrieval methods
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If vector_type is None or unsupported
|
||||||
|
"""
|
||||||
|
if vector_type is None:
|
||||||
|
raise ValueError("Vector store type is not configured.")
|
||||||
|
|
||||||
|
# Define vector database types that only support semantic search
|
||||||
|
semantic_only_types = {
|
||||||
|
VectorType.RELYT,
|
||||||
|
VectorType.TIDB_VECTOR,
|
||||||
|
VectorType.CHROMA,
|
||||||
|
VectorType.PGVECTO_RS,
|
||||||
|
VectorType.VIKINGDB,
|
||||||
|
VectorType.UPSTASH,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Define vector database types that support all retrieval methods
|
||||||
|
full_search_types = {
|
||||||
|
VectorType.QDRANT,
|
||||||
|
VectorType.WEAVIATE,
|
||||||
|
VectorType.OPENSEARCH,
|
||||||
|
VectorType.ANALYTICDB,
|
||||||
|
VectorType.MYSCALE,
|
||||||
|
VectorType.ORACLE,
|
||||||
|
VectorType.ELASTICSEARCH,
|
||||||
|
VectorType.ELASTICSEARCH_JA,
|
||||||
|
VectorType.PGVECTOR,
|
||||||
|
VectorType.VASTBASE,
|
||||||
|
VectorType.TIDB_ON_QDRANT,
|
||||||
|
VectorType.LINDORM,
|
||||||
|
VectorType.COUCHBASE,
|
||||||
|
VectorType.OPENGAUSS,
|
||||||
|
VectorType.OCEANBASE,
|
||||||
|
VectorType.TABLESTORE,
|
||||||
|
VectorType.HUAWEI_CLOUD,
|
||||||
|
VectorType.TENCENT,
|
||||||
|
VectorType.MATRIXONE,
|
||||||
|
VectorType.CLICKZETTA,
|
||||||
|
VectorType.BAIDU,
|
||||||
|
VectorType.ALIBABACLOUD_MYSQL,
|
||||||
|
}
|
||||||
|
|
||||||
|
semantic_methods = {"retrieval_method": [RetrievalMethod.SEMANTIC_SEARCH.value]}
|
||||||
|
full_methods = {
|
||||||
|
"retrieval_method": [
|
||||||
|
RetrievalMethod.SEMANTIC_SEARCH.value,
|
||||||
|
RetrievalMethod.FULL_TEXT_SEARCH.value,
|
||||||
|
RetrievalMethod.HYBRID_SEARCH.value,
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
if vector_type == VectorType.MILVUS:
|
||||||
|
return semantic_methods if is_mock else full_methods
|
||||||
|
|
||||||
|
if vector_type in semantic_only_types:
|
||||||
|
return semantic_methods
|
||||||
|
elif vector_type in full_search_types:
|
||||||
|
return full_methods
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unsupported vector db type {vector_type}.")
|
||||||
|
|
||||||
|
|
||||||
@console_ns.route("/datasets")
|
@console_ns.route("/datasets")
|
||||||
class DatasetListApi(Resource):
|
class DatasetListApi(Resource):
|
||||||
@api.doc("get_datasets")
|
@api.doc("get_datasets")
|
||||||
@ -500,7 +573,7 @@ class DatasetIndexingEstimateApi(Resource):
|
|||||||
if file_details:
|
if file_details:
|
||||||
for file_detail in file_details:
|
for file_detail in file_details:
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.FILE.value,
|
datasource_type=DatasourceType.FILE,
|
||||||
upload_file=file_detail,
|
upload_file=file_detail,
|
||||||
document_model=args["doc_form"],
|
document_model=args["doc_form"],
|
||||||
)
|
)
|
||||||
@ -512,14 +585,16 @@ class DatasetIndexingEstimateApi(Resource):
|
|||||||
credential_id = notion_info.get("credential_id")
|
credential_id = notion_info.get("credential_id")
|
||||||
for page in notion_info["pages"]:
|
for page in notion_info["pages"]:
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.NOTION.value,
|
datasource_type=DatasourceType.NOTION,
|
||||||
notion_info={
|
notion_info=NotionInfo.model_validate(
|
||||||
"credential_id": credential_id,
|
{
|
||||||
"notion_workspace_id": workspace_id,
|
"credential_id": credential_id,
|
||||||
"notion_obj_id": page["page_id"],
|
"notion_workspace_id": workspace_id,
|
||||||
"notion_page_type": page["type"],
|
"notion_obj_id": page["page_id"],
|
||||||
"tenant_id": current_user.current_tenant_id,
|
"notion_page_type": page["type"],
|
||||||
},
|
"tenant_id": current_user.current_tenant_id,
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=args["doc_form"],
|
document_model=args["doc_form"],
|
||||||
)
|
)
|
||||||
extract_settings.append(extract_setting)
|
extract_settings.append(extract_setting)
|
||||||
@ -527,15 +602,17 @@ class DatasetIndexingEstimateApi(Resource):
|
|||||||
website_info_list = args["info_list"]["website_info_list"]
|
website_info_list = args["info_list"]["website_info_list"]
|
||||||
for url in website_info_list["urls"]:
|
for url in website_info_list["urls"]:
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.WEBSITE.value,
|
datasource_type=DatasourceType.WEBSITE,
|
||||||
website_info={
|
website_info=WebsiteInfo.model_validate(
|
||||||
"provider": website_info_list["provider"],
|
{
|
||||||
"job_id": website_info_list["job_id"],
|
"provider": website_info_list["provider"],
|
||||||
"url": url,
|
"job_id": website_info_list["job_id"],
|
||||||
"tenant_id": current_user.current_tenant_id,
|
"url": url,
|
||||||
"mode": "crawl",
|
"tenant_id": current_user.current_tenant_id,
|
||||||
"only_main_content": website_info_list["only_main_content"],
|
"mode": "crawl",
|
||||||
},
|
"only_main_content": website_info_list["only_main_content"],
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=args["doc_form"],
|
document_model=args["doc_form"],
|
||||||
)
|
)
|
||||||
extract_settings.append(extract_setting)
|
extract_settings.append(extract_setting)
|
||||||
@ -773,49 +850,7 @@ class DatasetRetrievalSettingApi(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self):
|
def get(self):
|
||||||
vector_type = dify_config.VECTOR_STORE
|
vector_type = dify_config.VECTOR_STORE
|
||||||
match vector_type:
|
return _get_retrieval_methods_by_vector_type(vector_type, is_mock=False)
|
||||||
case (
|
|
||||||
VectorType.RELYT
|
|
||||||
| VectorType.TIDB_VECTOR
|
|
||||||
| VectorType.CHROMA
|
|
||||||
| VectorType.PGVECTO_RS
|
|
||||||
| VectorType.VIKINGDB
|
|
||||||
| VectorType.UPSTASH
|
|
||||||
):
|
|
||||||
return {"retrieval_method": [RetrievalMethod.SEMANTIC_SEARCH.value]}
|
|
||||||
case (
|
|
||||||
VectorType.QDRANT
|
|
||||||
| VectorType.WEAVIATE
|
|
||||||
| VectorType.OPENSEARCH
|
|
||||||
| VectorType.ANALYTICDB
|
|
||||||
| VectorType.MYSCALE
|
|
||||||
| VectorType.ORACLE
|
|
||||||
| VectorType.ELASTICSEARCH
|
|
||||||
| VectorType.ELASTICSEARCH_JA
|
|
||||||
| VectorType.PGVECTOR
|
|
||||||
| VectorType.VASTBASE
|
|
||||||
| VectorType.TIDB_ON_QDRANT
|
|
||||||
| VectorType.LINDORM
|
|
||||||
| VectorType.COUCHBASE
|
|
||||||
| VectorType.MILVUS
|
|
||||||
| VectorType.OPENGAUSS
|
|
||||||
| VectorType.OCEANBASE
|
|
||||||
| VectorType.TABLESTORE
|
|
||||||
| VectorType.HUAWEI_CLOUD
|
|
||||||
| VectorType.TENCENT
|
|
||||||
| VectorType.MATRIXONE
|
|
||||||
| VectorType.CLICKZETTA
|
|
||||||
| VectorType.BAIDU
|
|
||||||
):
|
|
||||||
return {
|
|
||||||
"retrieval_method": [
|
|
||||||
RetrievalMethod.SEMANTIC_SEARCH.value,
|
|
||||||
RetrievalMethod.FULL_TEXT_SEARCH.value,
|
|
||||||
RetrievalMethod.HYBRID_SEARCH.value,
|
|
||||||
]
|
|
||||||
}
|
|
||||||
case _:
|
|
||||||
raise ValueError(f"Unsupported vector db type {vector_type}.")
|
|
||||||
|
|
||||||
|
|
||||||
@console_ns.route("/datasets/retrieval-setting/<string:vector_type>")
|
@console_ns.route("/datasets/retrieval-setting/<string:vector_type>")
|
||||||
@ -828,48 +863,7 @@ class DatasetRetrievalSettingMockApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self, vector_type):
|
def get(self, vector_type):
|
||||||
match vector_type:
|
return _get_retrieval_methods_by_vector_type(vector_type, is_mock=True)
|
||||||
case (
|
|
||||||
VectorType.MILVUS
|
|
||||||
| VectorType.RELYT
|
|
||||||
| VectorType.TIDB_VECTOR
|
|
||||||
| VectorType.CHROMA
|
|
||||||
| VectorType.PGVECTO_RS
|
|
||||||
| VectorType.VIKINGDB
|
|
||||||
| VectorType.UPSTASH
|
|
||||||
):
|
|
||||||
return {"retrieval_method": [RetrievalMethod.SEMANTIC_SEARCH.value]}
|
|
||||||
case (
|
|
||||||
VectorType.QDRANT
|
|
||||||
| VectorType.WEAVIATE
|
|
||||||
| VectorType.OPENSEARCH
|
|
||||||
| VectorType.ANALYTICDB
|
|
||||||
| VectorType.MYSCALE
|
|
||||||
| VectorType.ORACLE
|
|
||||||
| VectorType.ELASTICSEARCH
|
|
||||||
| VectorType.ELASTICSEARCH_JA
|
|
||||||
| VectorType.COUCHBASE
|
|
||||||
| VectorType.PGVECTOR
|
|
||||||
| VectorType.VASTBASE
|
|
||||||
| VectorType.LINDORM
|
|
||||||
| VectorType.OPENGAUSS
|
|
||||||
| VectorType.OCEANBASE
|
|
||||||
| VectorType.TABLESTORE
|
|
||||||
| VectorType.TENCENT
|
|
||||||
| VectorType.HUAWEI_CLOUD
|
|
||||||
| VectorType.MATRIXONE
|
|
||||||
| VectorType.CLICKZETTA
|
|
||||||
| VectorType.BAIDU
|
|
||||||
):
|
|
||||||
return {
|
|
||||||
"retrieval_method": [
|
|
||||||
RetrievalMethod.SEMANTIC_SEARCH.value,
|
|
||||||
RetrievalMethod.FULL_TEXT_SEARCH.value,
|
|
||||||
RetrievalMethod.HYBRID_SEARCH.value,
|
|
||||||
]
|
|
||||||
}
|
|
||||||
case _:
|
|
||||||
raise ValueError(f"Unsupported vector db type {vector_type}.")
|
|
||||||
|
|
||||||
|
|
||||||
@console_ns.route("/datasets/<uuid:dataset_id>/error-docs")
|
@console_ns.route("/datasets/<uuid:dataset_id>/error-docs")
|
||||||
|
|||||||
@ -44,7 +44,7 @@ from core.model_runtime.entities.model_entities import ModelType
|
|||||||
from core.model_runtime.errors.invoke import InvokeAuthorizationError
|
from core.model_runtime.errors.invoke import InvokeAuthorizationError
|
||||||
from core.plugin.impl.exc import PluginDaemonClientSideError
|
from core.plugin.impl.exc import PluginDaemonClientSideError
|
||||||
from core.rag.extractor.entity.datasource_type import DatasourceType
|
from core.rag.extractor.entity.datasource_type import DatasourceType
|
||||||
from core.rag.extractor.entity.extract_setting import ExtractSetting
|
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from fields.document_fields import (
|
from fields.document_fields import (
|
||||||
dataset_and_document_fields,
|
dataset_and_document_fields,
|
||||||
@ -305,7 +305,7 @@ class DatasetDocumentListApi(Resource):
|
|||||||
"doc_language", type=str, default="English", required=False, nullable=False, location="json"
|
"doc_language", type=str, default="English", required=False, nullable=False, location="json"
|
||||||
)
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
knowledge_config = KnowledgeConfig(**args)
|
knowledge_config = KnowledgeConfig.model_validate(args)
|
||||||
|
|
||||||
if not dataset.indexing_technique and not knowledge_config.indexing_technique:
|
if not dataset.indexing_technique and not knowledge_config.indexing_technique:
|
||||||
raise ValueError("indexing_technique is required.")
|
raise ValueError("indexing_technique is required.")
|
||||||
@ -395,7 +395,7 @@ class DatasetInitApi(Resource):
|
|||||||
parser.add_argument("embedding_model_provider", type=str, required=False, nullable=True, location="json")
|
parser.add_argument("embedding_model_provider", type=str, required=False, nullable=True, location="json")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
knowledge_config = KnowledgeConfig(**args)
|
knowledge_config = KnowledgeConfig.model_validate(args)
|
||||||
if knowledge_config.indexing_technique == "high_quality":
|
if knowledge_config.indexing_technique == "high_quality":
|
||||||
if knowledge_config.embedding_model is None or knowledge_config.embedding_model_provider is None:
|
if knowledge_config.embedding_model is None or knowledge_config.embedding_model_provider is None:
|
||||||
raise ValueError("embedding model and embedding model provider are required for high quality indexing.")
|
raise ValueError("embedding model and embedding model provider are required for high quality indexing.")
|
||||||
@ -475,7 +475,7 @@ class DocumentIndexingEstimateApi(DocumentResource):
|
|||||||
raise NotFound("File not found.")
|
raise NotFound("File not found.")
|
||||||
|
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.FILE.value, upload_file=file, document_model=document.doc_form
|
datasource_type=DatasourceType.FILE, upload_file=file, document_model=document.doc_form
|
||||||
)
|
)
|
||||||
|
|
||||||
indexing_runner = IndexingRunner()
|
indexing_runner = IndexingRunner()
|
||||||
@ -538,7 +538,7 @@ class DocumentBatchIndexingEstimateApi(DocumentResource):
|
|||||||
raise NotFound("File not found.")
|
raise NotFound("File not found.")
|
||||||
|
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.FILE.value, upload_file=file_detail, document_model=document.doc_form
|
datasource_type=DatasourceType.FILE, upload_file=file_detail, document_model=document.doc_form
|
||||||
)
|
)
|
||||||
extract_settings.append(extract_setting)
|
extract_settings.append(extract_setting)
|
||||||
|
|
||||||
@ -546,14 +546,16 @@ class DocumentBatchIndexingEstimateApi(DocumentResource):
|
|||||||
if not data_source_info:
|
if not data_source_info:
|
||||||
continue
|
continue
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.NOTION.value,
|
datasource_type=DatasourceType.NOTION,
|
||||||
notion_info={
|
notion_info=NotionInfo.model_validate(
|
||||||
"credential_id": data_source_info["credential_id"],
|
{
|
||||||
"notion_workspace_id": data_source_info["notion_workspace_id"],
|
"credential_id": data_source_info["credential_id"],
|
||||||
"notion_obj_id": data_source_info["notion_page_id"],
|
"notion_workspace_id": data_source_info["notion_workspace_id"],
|
||||||
"notion_page_type": data_source_info["type"],
|
"notion_obj_id": data_source_info["notion_page_id"],
|
||||||
"tenant_id": current_user.current_tenant_id,
|
"notion_page_type": data_source_info["type"],
|
||||||
},
|
"tenant_id": current_user.current_tenant_id,
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=document.doc_form,
|
document_model=document.doc_form,
|
||||||
)
|
)
|
||||||
extract_settings.append(extract_setting)
|
extract_settings.append(extract_setting)
|
||||||
@ -561,15 +563,17 @@ class DocumentBatchIndexingEstimateApi(DocumentResource):
|
|||||||
if not data_source_info:
|
if not data_source_info:
|
||||||
continue
|
continue
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.WEBSITE.value,
|
datasource_type=DatasourceType.WEBSITE,
|
||||||
website_info={
|
website_info=WebsiteInfo.model_validate(
|
||||||
"provider": data_source_info["provider"],
|
{
|
||||||
"job_id": data_source_info["job_id"],
|
"provider": data_source_info["provider"],
|
||||||
"url": data_source_info["url"],
|
"job_id": data_source_info["job_id"],
|
||||||
"tenant_id": current_user.current_tenant_id,
|
"url": data_source_info["url"],
|
||||||
"mode": data_source_info["mode"],
|
"tenant_id": current_user.current_tenant_id,
|
||||||
"only_main_content": data_source_info["only_main_content"],
|
"mode": data_source_info["mode"],
|
||||||
},
|
"only_main_content": data_source_info["only_main_content"],
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=document.doc_form,
|
document_model=document.doc_form,
|
||||||
)
|
)
|
||||||
extract_settings.append(extract_setting)
|
extract_settings.append(extract_setting)
|
||||||
|
|||||||
@ -309,7 +309,7 @@ class DatasetDocumentSegmentUpdateApi(Resource):
|
|||||||
)
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
SegmentService.segment_create_args_validate(args, document)
|
SegmentService.segment_create_args_validate(args, document)
|
||||||
segment = SegmentService.update_segment(SegmentUpdateArgs(**args), segment, document, dataset)
|
segment = SegmentService.update_segment(SegmentUpdateArgs.model_validate(args), segment, document, dataset)
|
||||||
return {"data": marshal(segment, segment_fields), "doc_form": document.doc_form}, 200
|
return {"data": marshal(segment, segment_fields), "doc_form": document.doc_form}, 200
|
||||||
|
|
||||||
@setup_required
|
@setup_required
|
||||||
@ -564,7 +564,7 @@ class ChildChunkAddApi(Resource):
|
|||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
try:
|
try:
|
||||||
chunks_data = args["chunks"]
|
chunks_data = args["chunks"]
|
||||||
chunks = [ChildChunkUpdateArgs(**chunk) for chunk in chunks_data]
|
chunks = [ChildChunkUpdateArgs.model_validate(chunk) for chunk in chunks_data]
|
||||||
child_chunks = SegmentService.update_child_chunks(chunks, segment, document, dataset)
|
child_chunks = SegmentService.update_child_chunks(chunks, segment, document, dataset)
|
||||||
except ChildChunkIndexingServiceError as e:
|
except ChildChunkIndexingServiceError as e:
|
||||||
raise ChildChunkIndexingError(str(e))
|
raise ChildChunkIndexingError(str(e))
|
||||||
|
|||||||
@ -1,7 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
from typing import cast
|
|
||||||
|
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import marshal, reqparse
|
from flask_restx import marshal, reqparse
|
||||||
from werkzeug.exceptions import Forbidden, InternalServerError, NotFound
|
from werkzeug.exceptions import Forbidden, InternalServerError, NotFound
|
||||||
|
|
||||||
@ -21,6 +19,7 @@ from core.errors.error import (
|
|||||||
)
|
)
|
||||||
from core.model_runtime.errors.invoke import InvokeError
|
from core.model_runtime.errors.invoke import InvokeError
|
||||||
from fields.hit_testing_fields import hit_testing_record_fields
|
from fields.hit_testing_fields import hit_testing_record_fields
|
||||||
|
from libs.login import current_user
|
||||||
from models.account import Account
|
from models.account import Account
|
||||||
from services.dataset_service import DatasetService
|
from services.dataset_service import DatasetService
|
||||||
from services.hit_testing_service import HitTestingService
|
from services.hit_testing_service import HitTestingService
|
||||||
@ -31,6 +30,7 @@ logger = logging.getLogger(__name__)
|
|||||||
class DatasetsHitTestingBase:
|
class DatasetsHitTestingBase:
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def get_and_validate_dataset(dataset_id: str):
|
def get_and_validate_dataset(dataset_id: str):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
dataset = DatasetService.get_dataset(dataset_id)
|
dataset = DatasetService.get_dataset(dataset_id)
|
||||||
if dataset is None:
|
if dataset is None:
|
||||||
raise NotFound("Dataset not found.")
|
raise NotFound("Dataset not found.")
|
||||||
@ -57,11 +57,12 @@ class DatasetsHitTestingBase:
|
|||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def perform_hit_testing(dataset, args):
|
def perform_hit_testing(dataset, args):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
try:
|
try:
|
||||||
response = HitTestingService.retrieve(
|
response = HitTestingService.retrieve(
|
||||||
dataset=dataset,
|
dataset=dataset,
|
||||||
query=args["query"],
|
query=args["query"],
|
||||||
account=cast(Account, current_user),
|
account=current_user,
|
||||||
retrieval_model=args["retrieval_model"],
|
retrieval_model=args["retrieval_model"],
|
||||||
external_retrieval_model=args["external_retrieval_model"],
|
external_retrieval_model=args["external_retrieval_model"],
|
||||||
limit=10,
|
limit=10,
|
||||||
|
|||||||
@ -28,7 +28,7 @@ class DatasetMetadataCreateApi(Resource):
|
|||||||
parser.add_argument("type", type=str, required=True, nullable=False, location="json")
|
parser.add_argument("type", type=str, required=True, nullable=False, location="json")
|
||||||
parser.add_argument("name", type=str, required=True, nullable=False, location="json")
|
parser.add_argument("name", type=str, required=True, nullable=False, location="json")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
metadata_args = MetadataArgs(**args)
|
metadata_args = MetadataArgs.model_validate(args)
|
||||||
|
|
||||||
dataset_id_str = str(dataset_id)
|
dataset_id_str = str(dataset_id)
|
||||||
dataset = DatasetService.get_dataset(dataset_id_str)
|
dataset = DatasetService.get_dataset(dataset_id_str)
|
||||||
@ -137,7 +137,7 @@ class DocumentMetadataEditApi(Resource):
|
|||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("operation_data", type=list, required=True, nullable=False, location="json")
|
parser.add_argument("operation_data", type=list, required=True, nullable=False, location="json")
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
metadata_args = MetadataOperationData(**args)
|
metadata_args = MetadataOperationData.model_validate(args)
|
||||||
|
|
||||||
MetadataService.update_documents_metadata(dataset, metadata_args)
|
MetadataService.update_documents_metadata(dataset, metadata_args)
|
||||||
|
|
||||||
|
|||||||
@ -88,7 +88,7 @@ class CustomizedPipelineTemplateApi(Resource):
|
|||||||
nullable=True,
|
nullable=True,
|
||||||
)
|
)
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
pipeline_template_info = PipelineTemplateInfoEntity(**args)
|
pipeline_template_info = PipelineTemplateInfoEntity.model_validate(args)
|
||||||
RagPipelineService.update_customized_pipeline_template(template_id, pipeline_template_info)
|
RagPipelineService.update_customized_pipeline_template(template_id, pipeline_template_info)
|
||||||
return 200
|
return 200
|
||||||
|
|
||||||
|
|||||||
@ -60,9 +60,9 @@ class RagPipelineImportApi(Resource):
|
|||||||
|
|
||||||
# Return appropriate status code based on result
|
# Return appropriate status code based on result
|
||||||
status = result.status
|
status = result.status
|
||||||
if status == ImportStatus.FAILED.value:
|
if status == ImportStatus.FAILED:
|
||||||
return result.model_dump(mode="json"), 400
|
return result.model_dump(mode="json"), 400
|
||||||
elif status == ImportStatus.PENDING.value:
|
elif status == ImportStatus.PENDING:
|
||||||
return result.model_dump(mode="json"), 202
|
return result.model_dump(mode="json"), 202
|
||||||
return result.model_dump(mode="json"), 200
|
return result.model_dump(mode="json"), 200
|
||||||
|
|
||||||
@ -87,7 +87,7 @@ class RagPipelineImportConfirmApi(Resource):
|
|||||||
session.commit()
|
session.commit()
|
||||||
|
|
||||||
# Return appropriate status code based on result
|
# Return appropriate status code based on result
|
||||||
if result.status == ImportStatus.FAILED.value:
|
if result.status == ImportStatus.FAILED:
|
||||||
return result.model_dump(mode="json"), 400
|
return result.model_dump(mode="json"), 400
|
||||||
return result.model_dump(mode="json"), 200
|
return result.model_dump(mode="json"), 200
|
||||||
|
|
||||||
|
|||||||
@ -6,7 +6,7 @@ from flask_restx import Resource, inputs, marshal_with, reqparse
|
|||||||
from sqlalchemy import and_, select
|
from sqlalchemy import and_, select
|
||||||
from werkzeug.exceptions import BadRequest, Forbidden, NotFound
|
from werkzeug.exceptions import BadRequest, Forbidden, NotFound
|
||||||
|
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.explore.wraps import InstalledAppResource
|
from controllers.console.explore.wraps import InstalledAppResource
|
||||||
from controllers.console.wraps import account_initialization_required, cloud_edition_billing_resource_check
|
from controllers.console.wraps import account_initialization_required, cloud_edition_billing_resource_check
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
@ -22,6 +22,7 @@ from services.feature_service import FeatureService
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/installed-apps")
|
||||||
class InstalledAppsListApi(Resource):
|
class InstalledAppsListApi(Resource):
|
||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@ -154,6 +155,7 @@ class InstalledAppsListApi(Resource):
|
|||||||
return {"message": "App installed successfully"}
|
return {"message": "App installed successfully"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/installed-apps/<uuid:installed_app_id>")
|
||||||
class InstalledAppApi(InstalledAppResource):
|
class InstalledAppApi(InstalledAppResource):
|
||||||
"""
|
"""
|
||||||
update and delete an installed app
|
update and delete an installed app
|
||||||
@ -185,7 +187,3 @@ class InstalledAppApi(InstalledAppResource):
|
|||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
return {"result": "success", "message": "App info updated successfully"}
|
return {"result": "success", "message": "App info updated successfully"}
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(InstalledAppsListApi, "/installed-apps")
|
|
||||||
api.add_resource(InstalledAppApi, "/installed-apps/<uuid:installed_app_id>")
|
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
from flask_restx import marshal_with
|
from flask_restx import marshal_with
|
||||||
|
|
||||||
from controllers.common import fields
|
from controllers.common import fields
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.app.error import AppUnavailableError
|
from controllers.console.app.error import AppUnavailableError
|
||||||
from controllers.console.explore.wraps import InstalledAppResource
|
from controllers.console.explore.wraps import InstalledAppResource
|
||||||
from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict
|
from core.app.app_config.common.parameters_mapping import get_parameters_from_feature_dict
|
||||||
@ -9,6 +9,7 @@ from models.model import AppMode, InstalledApp
|
|||||||
from services.app_service import AppService
|
from services.app_service import AppService
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/installed-apps/<uuid:installed_app_id>/parameters", endpoint="installed_app_parameters")
|
||||||
class AppParameterApi(InstalledAppResource):
|
class AppParameterApi(InstalledAppResource):
|
||||||
"""Resource for app variables."""
|
"""Resource for app variables."""
|
||||||
|
|
||||||
@ -39,6 +40,7 @@ class AppParameterApi(InstalledAppResource):
|
|||||||
return get_parameters_from_feature_dict(features_dict=features_dict, user_input_form=user_input_form)
|
return get_parameters_from_feature_dict(features_dict=features_dict, user_input_form=user_input_form)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/installed-apps/<uuid:installed_app_id>/meta", endpoint="installed_app_meta")
|
||||||
class ExploreAppMetaApi(InstalledAppResource):
|
class ExploreAppMetaApi(InstalledAppResource):
|
||||||
def get(self, installed_app: InstalledApp):
|
def get(self, installed_app: InstalledApp):
|
||||||
"""Get app meta"""
|
"""Get app meta"""
|
||||||
@ -46,9 +48,3 @@ class ExploreAppMetaApi(InstalledAppResource):
|
|||||||
if not app_model:
|
if not app_model:
|
||||||
raise ValueError("App not found")
|
raise ValueError("App not found")
|
||||||
return AppService().get_app_meta(app_model)
|
return AppService().get_app_meta(app_model)
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(
|
|
||||||
AppParameterApi, "/installed-apps/<uuid:installed_app_id>/parameters", endpoint="installed_app_parameters"
|
|
||||||
)
|
|
||||||
api.add_resource(ExploreAppMetaApi, "/installed-apps/<uuid:installed_app_id>/meta", endpoint="installed_app_meta")
|
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
from flask_restx import Resource, fields, marshal_with, reqparse
|
from flask_restx import Resource, fields, marshal_with, reqparse
|
||||||
|
|
||||||
from constants.languages import languages
|
from constants.languages import languages
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.wraps import account_initialization_required
|
from controllers.console.wraps import account_initialization_required
|
||||||
from libs.helper import AppIconUrlField
|
from libs.helper import AppIconUrlField
|
||||||
from libs.login import current_user, login_required
|
from libs.login import current_user, login_required
|
||||||
@ -35,6 +35,7 @@ recommended_app_list_fields = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/explore/apps")
|
||||||
class RecommendedAppListApi(Resource):
|
class RecommendedAppListApi(Resource):
|
||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@ -56,13 +57,10 @@ class RecommendedAppListApi(Resource):
|
|||||||
return RecommendedAppService.get_recommended_apps_and_categories(language_prefix)
|
return RecommendedAppService.get_recommended_apps_and_categories(language_prefix)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/explore/apps/<uuid:app_id>")
|
||||||
class RecommendedAppApi(Resource):
|
class RecommendedAppApi(Resource):
|
||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self, app_id):
|
def get(self, app_id):
|
||||||
app_id = str(app_id)
|
app_id = str(app_id)
|
||||||
return RecommendedAppService.get_recommend_app_detail(app_id)
|
return RecommendedAppService.get_recommend_app_detail(app_id)
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(RecommendedAppListApi, "/explore/apps")
|
|
||||||
api.add_resource(RecommendedAppApi, "/explore/apps/<uuid:app_id>")
|
|
||||||
|
|||||||
@ -2,7 +2,7 @@ from flask_restx import fields, marshal_with, reqparse
|
|||||||
from flask_restx.inputs import int_range
|
from flask_restx.inputs import int_range
|
||||||
from werkzeug.exceptions import NotFound
|
from werkzeug.exceptions import NotFound
|
||||||
|
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.explore.error import NotCompletionAppError
|
from controllers.console.explore.error import NotCompletionAppError
|
||||||
from controllers.console.explore.wraps import InstalledAppResource
|
from controllers.console.explore.wraps import InstalledAppResource
|
||||||
from fields.conversation_fields import message_file_fields
|
from fields.conversation_fields import message_file_fields
|
||||||
@ -25,6 +25,7 @@ message_fields = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/installed-apps/<uuid:installed_app_id>/saved-messages", endpoint="installed_app_saved_messages")
|
||||||
class SavedMessageListApi(InstalledAppResource):
|
class SavedMessageListApi(InstalledAppResource):
|
||||||
saved_message_infinite_scroll_pagination_fields = {
|
saved_message_infinite_scroll_pagination_fields = {
|
||||||
"limit": fields.Integer,
|
"limit": fields.Integer,
|
||||||
@ -66,6 +67,9 @@ class SavedMessageListApi(InstalledAppResource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route(
|
||||||
|
"/installed-apps/<uuid:installed_app_id>/saved-messages/<uuid:message_id>", endpoint="installed_app_saved_message"
|
||||||
|
)
|
||||||
class SavedMessageApi(InstalledAppResource):
|
class SavedMessageApi(InstalledAppResource):
|
||||||
def delete(self, installed_app, message_id):
|
def delete(self, installed_app, message_id):
|
||||||
app_model = installed_app.app
|
app_model = installed_app.app
|
||||||
@ -80,15 +84,3 @@ class SavedMessageApi(InstalledAppResource):
|
|||||||
SavedMessageService.delete(app_model, current_user, message_id)
|
SavedMessageService.delete(app_model, current_user, message_id)
|
||||||
|
|
||||||
return {"result": "success"}, 204
|
return {"result": "success"}, 204
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(
|
|
||||||
SavedMessageListApi,
|
|
||||||
"/installed-apps/<uuid:installed_app_id>/saved-messages",
|
|
||||||
endpoint="installed_app_saved_messages",
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
SavedMessageApi,
|
|
||||||
"/installed-apps/<uuid:installed_app_id>/saved-messages/<uuid:message_id>",
|
|
||||||
endpoint="installed_app_saved_message",
|
|
||||||
)
|
|
||||||
|
|||||||
@ -2,15 +2,15 @@ from collections.abc import Callable
|
|||||||
from functools import wraps
|
from functools import wraps
|
||||||
from typing import Concatenate, ParamSpec, TypeVar
|
from typing import Concatenate, ParamSpec, TypeVar
|
||||||
|
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource
|
from flask_restx import Resource
|
||||||
from werkzeug.exceptions import NotFound
|
from werkzeug.exceptions import NotFound
|
||||||
|
|
||||||
from controllers.console.explore.error import AppAccessDeniedError
|
from controllers.console.explore.error import AppAccessDeniedError
|
||||||
from controllers.console.wraps import account_initialization_required
|
from controllers.console.wraps import account_initialization_required
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
from models import InstalledApp
|
from models import InstalledApp
|
||||||
|
from models.account import Account
|
||||||
from services.app_service import AppService
|
from services.app_service import AppService
|
||||||
from services.enterprise.enterprise_service import EnterpriseService
|
from services.enterprise.enterprise_service import EnterpriseService
|
||||||
from services.feature_service import FeatureService
|
from services.feature_service import FeatureService
|
||||||
@ -24,6 +24,8 @@ def installed_app_required(view: Callable[Concatenate[InstalledApp, P], R] | Non
|
|||||||
def decorator(view: Callable[Concatenate[InstalledApp, P], R]):
|
def decorator(view: Callable[Concatenate[InstalledApp, P], R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(installed_app_id: str, *args: P.args, **kwargs: P.kwargs):
|
def decorated(installed_app_id: str, *args: P.args, **kwargs: P.kwargs):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
installed_app = (
|
installed_app = (
|
||||||
db.session.query(InstalledApp)
|
db.session.query(InstalledApp)
|
||||||
.where(
|
.where(
|
||||||
@ -56,6 +58,7 @@ def user_allowed_to_access_app(view: Callable[Concatenate[InstalledApp, P], R] |
|
|||||||
def decorated(installed_app: InstalledApp, *args: P.args, **kwargs: P.kwargs):
|
def decorated(installed_app: InstalledApp, *args: P.args, **kwargs: P.kwargs):
|
||||||
feature = FeatureService.get_system_features()
|
feature = FeatureService.get_system_features()
|
||||||
if feature.webapp_auth.enabled:
|
if feature.webapp_auth.enabled:
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
app_id = installed_app.app_id
|
app_id = installed_app.app_id
|
||||||
app_code = AppService.get_app_code_by_id(app_id)
|
app_code = AppService.get_app_code_by_id(app_id)
|
||||||
res = EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp(
|
res = EnterpriseService.WebAppAuth.is_user_allowed_to_access_webapp(
|
||||||
|
|||||||
@ -1,11 +1,11 @@
|
|||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, fields, marshal_with, reqparse
|
from flask_restx import Resource, fields, marshal_with, reqparse
|
||||||
|
|
||||||
from constants import HIDDEN_VALUE
|
from constants import HIDDEN_VALUE
|
||||||
from controllers.console import api, console_ns
|
from controllers.console import api, console_ns
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from fields.api_based_extension_fields import api_based_extension_fields
|
from fields.api_based_extension_fields import api_based_extension_fields
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from models.api_based_extension import APIBasedExtension
|
from models.api_based_extension import APIBasedExtension
|
||||||
from services.api_based_extension_service import APIBasedExtensionService
|
from services.api_based_extension_service import APIBasedExtensionService
|
||||||
from services.code_based_extension_service import CodeBasedExtensionService
|
from services.code_based_extension_service import CodeBasedExtensionService
|
||||||
@ -47,6 +47,8 @@ class APIBasedExtensionAPI(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@marshal_with(api_based_extension_fields)
|
@marshal_with(api_based_extension_fields)
|
||||||
def get(self):
|
def get(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
tenant_id = current_user.current_tenant_id
|
tenant_id = current_user.current_tenant_id
|
||||||
return APIBasedExtensionService.get_all_by_tenant_id(tenant_id)
|
return APIBasedExtensionService.get_all_by_tenant_id(tenant_id)
|
||||||
|
|
||||||
@ -68,6 +70,8 @@ class APIBasedExtensionAPI(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@marshal_with(api_based_extension_fields)
|
@marshal_with(api_based_extension_fields)
|
||||||
def post(self):
|
def post(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("name", type=str, required=True, location="json")
|
parser.add_argument("name", type=str, required=True, location="json")
|
||||||
parser.add_argument("api_endpoint", type=str, required=True, location="json")
|
parser.add_argument("api_endpoint", type=str, required=True, location="json")
|
||||||
@ -95,6 +99,8 @@ class APIBasedExtensionDetailAPI(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@marshal_with(api_based_extension_fields)
|
@marshal_with(api_based_extension_fields)
|
||||||
def get(self, id):
|
def get(self, id):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
api_based_extension_id = str(id)
|
api_based_extension_id = str(id)
|
||||||
tenant_id = current_user.current_tenant_id
|
tenant_id = current_user.current_tenant_id
|
||||||
|
|
||||||
@ -119,6 +125,8 @@ class APIBasedExtensionDetailAPI(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@marshal_with(api_based_extension_fields)
|
@marshal_with(api_based_extension_fields)
|
||||||
def post(self, id):
|
def post(self, id):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
api_based_extension_id = str(id)
|
api_based_extension_id = str(id)
|
||||||
tenant_id = current_user.current_tenant_id
|
tenant_id = current_user.current_tenant_id
|
||||||
|
|
||||||
@ -146,6 +154,8 @@ class APIBasedExtensionDetailAPI(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def delete(self, id):
|
def delete(self, id):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
api_based_extension_id = str(id)
|
api_based_extension_id = str(id)
|
||||||
tenant_id = current_user.current_tenant_id
|
tenant_id = current_user.current_tenant_id
|
||||||
|
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, fields
|
from flask_restx import Resource, fields
|
||||||
|
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from services.feature_service import FeatureService
|
from services.feature_service import FeatureService
|
||||||
|
|
||||||
from . import api, console_ns
|
from . import api, console_ns
|
||||||
@ -23,6 +23,8 @@ class FeatureApi(Resource):
|
|||||||
@cloud_utm_record
|
@cloud_utm_record
|
||||||
def get(self):
|
def get(self):
|
||||||
"""Get feature configuration for current tenant"""
|
"""Get feature configuration for current tenant"""
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
return FeatureService.get_features(current_user.current_tenant_id).model_dump()
|
return FeatureService.get_features(current_user.current_tenant_id).model_dump()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@ -1,8 +1,6 @@
|
|||||||
import urllib.parse
|
import urllib.parse
|
||||||
from typing import cast
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, marshal_with, reqparse
|
from flask_restx import Resource, marshal_with, reqparse
|
||||||
|
|
||||||
import services
|
import services
|
||||||
@ -16,6 +14,7 @@ from core.file import helpers as file_helpers
|
|||||||
from core.helper import ssrf_proxy
|
from core.helper import ssrf_proxy
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from fields.file_fields import file_fields_with_signed_url, remote_file_info_fields
|
from fields.file_fields import file_fields_with_signed_url, remote_file_info_fields
|
||||||
|
from libs.login import current_user
|
||||||
from models.account import Account
|
from models.account import Account
|
||||||
from services.file_service import FileService
|
from services.file_service import FileService
|
||||||
|
|
||||||
@ -65,7 +64,8 @@ class RemoteFileUploadApi(Resource):
|
|||||||
content = resp.content if resp.request.method == "GET" else ssrf_proxy.get(url).content
|
content = resp.content if resp.request.method == "GET" else ssrf_proxy.get(url).content
|
||||||
|
|
||||||
try:
|
try:
|
||||||
user = cast(Account, current_user)
|
assert isinstance(current_user, Account)
|
||||||
|
user = current_user
|
||||||
upload_file = FileService(db.engine).upload_file(
|
upload_file = FileService(db.engine).upload_file(
|
||||||
filename=file_info.filename,
|
filename=file_info.filename,
|
||||||
content=content,
|
content=content,
|
||||||
|
|||||||
@ -1,12 +1,12 @@
|
|||||||
from flask import request
|
from flask import request
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, marshal_with, reqparse
|
from flask_restx import Resource, marshal_with, reqparse
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from controllers.console import console_ns
|
from controllers.console import console_ns
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from fields.tag_fields import dataset_tag_fields
|
from fields.tag_fields import dataset_tag_fields
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from models.model import Tag
|
from models.model import Tag
|
||||||
from services.tag_service import TagService
|
from services.tag_service import TagService
|
||||||
|
|
||||||
@ -24,6 +24,8 @@ class TagListApi(Resource):
|
|||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
@marshal_with(dataset_tag_fields)
|
@marshal_with(dataset_tag_fields)
|
||||||
def get(self):
|
def get(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
tag_type = request.args.get("type", type=str, default="")
|
tag_type = request.args.get("type", type=str, default="")
|
||||||
keyword = request.args.get("keyword", default=None, type=str)
|
keyword = request.args.get("keyword", default=None, type=str)
|
||||||
tags = TagService.get_tags(tag_type, current_user.current_tenant_id, keyword)
|
tags = TagService.get_tags(tag_type, current_user.current_tenant_id, keyword)
|
||||||
@ -34,8 +36,10 @@ class TagListApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
# The role of the current user in the ta table must be admin, owner, or editor
|
# The role of the current user in the ta table must be admin, owner, or editor
|
||||||
if not (current_user.is_editor or current_user.is_dataset_editor):
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
@ -59,9 +63,11 @@ class TagUpdateDeleteApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def patch(self, tag_id):
|
def patch(self, tag_id):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
tag_id = str(tag_id)
|
tag_id = str(tag_id)
|
||||||
# The role of the current user in the ta table must be admin, owner, or editor
|
# The role of the current user in the ta table must be admin, owner, or editor
|
||||||
if not (current_user.is_editor or current_user.is_dataset_editor):
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
@ -81,9 +87,11 @@ class TagUpdateDeleteApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def delete(self, tag_id):
|
def delete(self, tag_id):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
tag_id = str(tag_id)
|
tag_id = str(tag_id)
|
||||||
# The role of the current user in the ta table must be admin, owner, or editor
|
# The role of the current user in the ta table must be admin, owner, or editor
|
||||||
if not current_user.is_editor:
|
if not current_user.has_edit_permission:
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
TagService.delete_tag(tag_id)
|
TagService.delete_tag(tag_id)
|
||||||
@ -97,8 +105,10 @@ class TagBindingCreateApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
||||||
if not (current_user.is_editor or current_user.is_dataset_editor):
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
@ -123,8 +133,10 @@ class TagBindingDeleteApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
assert current_user.current_tenant_id is not None
|
||||||
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
# The role of the current user in the ta table must be admin, owner, editor, or dataset_operator
|
||||||
if not (current_user.is_editor or current_user.is_dataset_editor):
|
if not (current_user.has_edit_permission or current_user.is_dataset_editor):
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
|
|||||||
@ -9,7 +9,7 @@ from sqlalchemy.orm import Session
|
|||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from constants.languages import supported_language
|
from constants.languages import supported_language
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.auth.error import (
|
from controllers.console.auth.error import (
|
||||||
EmailAlreadyInUseError,
|
EmailAlreadyInUseError,
|
||||||
EmailChangeLimitError,
|
EmailChangeLimitError,
|
||||||
@ -45,6 +45,7 @@ from services.billing_service import BillingService
|
|||||||
from services.errors.account import CurrentPasswordIncorrectError as ServiceCurrentPasswordIncorrectError
|
from services.errors.account import CurrentPasswordIncorrectError as ServiceCurrentPasswordIncorrectError
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/init")
|
||||||
class AccountInitApi(Resource):
|
class AccountInitApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -97,6 +98,7 @@ class AccountInitApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/profile")
|
||||||
class AccountProfileApi(Resource):
|
class AccountProfileApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -109,6 +111,7 @@ class AccountProfileApi(Resource):
|
|||||||
return current_user
|
return current_user
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/name")
|
||||||
class AccountNameApi(Resource):
|
class AccountNameApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -130,6 +133,7 @@ class AccountNameApi(Resource):
|
|||||||
return updated_account
|
return updated_account
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/avatar")
|
||||||
class AccountAvatarApi(Resource):
|
class AccountAvatarApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -147,6 +151,7 @@ class AccountAvatarApi(Resource):
|
|||||||
return updated_account
|
return updated_account
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/interface-language")
|
||||||
class AccountInterfaceLanguageApi(Resource):
|
class AccountInterfaceLanguageApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -164,6 +169,7 @@ class AccountInterfaceLanguageApi(Resource):
|
|||||||
return updated_account
|
return updated_account
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/interface-theme")
|
||||||
class AccountInterfaceThemeApi(Resource):
|
class AccountInterfaceThemeApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -181,6 +187,7 @@ class AccountInterfaceThemeApi(Resource):
|
|||||||
return updated_account
|
return updated_account
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/timezone")
|
||||||
class AccountTimezoneApi(Resource):
|
class AccountTimezoneApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -202,6 +209,7 @@ class AccountTimezoneApi(Resource):
|
|||||||
return updated_account
|
return updated_account
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/password")
|
||||||
class AccountPasswordApi(Resource):
|
class AccountPasswordApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -227,6 +235,7 @@ class AccountPasswordApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/integrates")
|
||||||
class AccountIntegrateApi(Resource):
|
class AccountIntegrateApi(Resource):
|
||||||
integrate_fields = {
|
integrate_fields = {
|
||||||
"provider": fields.String,
|
"provider": fields.String,
|
||||||
@ -283,6 +292,7 @@ class AccountIntegrateApi(Resource):
|
|||||||
return {"data": integrate_data}
|
return {"data": integrate_data}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/delete/verify")
|
||||||
class AccountDeleteVerifyApi(Resource):
|
class AccountDeleteVerifyApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -298,6 +308,7 @@ class AccountDeleteVerifyApi(Resource):
|
|||||||
return {"result": "success", "data": token}
|
return {"result": "success", "data": token}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/delete")
|
||||||
class AccountDeleteApi(Resource):
|
class AccountDeleteApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -320,6 +331,7 @@ class AccountDeleteApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/delete/feedback")
|
||||||
class AccountDeleteUpdateFeedbackApi(Resource):
|
class AccountDeleteUpdateFeedbackApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
def post(self):
|
def post(self):
|
||||||
@ -333,6 +345,7 @@ class AccountDeleteUpdateFeedbackApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/education/verify")
|
||||||
class EducationVerifyApi(Resource):
|
class EducationVerifyApi(Resource):
|
||||||
verify_fields = {
|
verify_fields = {
|
||||||
"token": fields.String,
|
"token": fields.String,
|
||||||
@ -352,6 +365,7 @@ class EducationVerifyApi(Resource):
|
|||||||
return BillingService.EducationIdentity.verify(account.id, account.email)
|
return BillingService.EducationIdentity.verify(account.id, account.email)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/education")
|
||||||
class EducationApi(Resource):
|
class EducationApi(Resource):
|
||||||
status_fields = {
|
status_fields = {
|
||||||
"result": fields.Boolean,
|
"result": fields.Boolean,
|
||||||
@ -396,6 +410,7 @@ class EducationApi(Resource):
|
|||||||
return res
|
return res
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/education/autocomplete")
|
||||||
class EducationAutoCompleteApi(Resource):
|
class EducationAutoCompleteApi(Resource):
|
||||||
data_fields = {
|
data_fields = {
|
||||||
"data": fields.List(fields.String),
|
"data": fields.List(fields.String),
|
||||||
@ -419,6 +434,7 @@ class EducationAutoCompleteApi(Resource):
|
|||||||
return BillingService.EducationIdentity.autocomplete(args["keywords"], args["page"], args["limit"])
|
return BillingService.EducationIdentity.autocomplete(args["keywords"], args["page"], args["limit"])
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/change-email")
|
||||||
class ChangeEmailSendEmailApi(Resource):
|
class ChangeEmailSendEmailApi(Resource):
|
||||||
@enable_change_email
|
@enable_change_email
|
||||||
@setup_required
|
@setup_required
|
||||||
@ -467,6 +483,7 @@ class ChangeEmailSendEmailApi(Resource):
|
|||||||
return {"result": "success", "data": token}
|
return {"result": "success", "data": token}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/change-email/validity")
|
||||||
class ChangeEmailCheckApi(Resource):
|
class ChangeEmailCheckApi(Resource):
|
||||||
@enable_change_email
|
@enable_change_email
|
||||||
@setup_required
|
@setup_required
|
||||||
@ -508,6 +525,7 @@ class ChangeEmailCheckApi(Resource):
|
|||||||
return {"is_valid": True, "email": token_data.get("email"), "token": new_token}
|
return {"is_valid": True, "email": token_data.get("email"), "token": new_token}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/change-email/reset")
|
||||||
class ChangeEmailResetApi(Resource):
|
class ChangeEmailResetApi(Resource):
|
||||||
@enable_change_email
|
@enable_change_email
|
||||||
@setup_required
|
@setup_required
|
||||||
@ -547,6 +565,7 @@ class ChangeEmailResetApi(Resource):
|
|||||||
return updated_account
|
return updated_account
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/account/change-email/check-email-unique")
|
||||||
class CheckEmailUnique(Resource):
|
class CheckEmailUnique(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
def post(self):
|
def post(self):
|
||||||
@ -558,28 +577,3 @@ class CheckEmailUnique(Resource):
|
|||||||
if not AccountService.check_email_unique(args["email"]):
|
if not AccountService.check_email_unique(args["email"]):
|
||||||
raise EmailAlreadyInUseError()
|
raise EmailAlreadyInUseError()
|
||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
# Register API resources
|
|
||||||
api.add_resource(AccountInitApi, "/account/init")
|
|
||||||
api.add_resource(AccountProfileApi, "/account/profile")
|
|
||||||
api.add_resource(AccountNameApi, "/account/name")
|
|
||||||
api.add_resource(AccountAvatarApi, "/account/avatar")
|
|
||||||
api.add_resource(AccountInterfaceLanguageApi, "/account/interface-language")
|
|
||||||
api.add_resource(AccountInterfaceThemeApi, "/account/interface-theme")
|
|
||||||
api.add_resource(AccountTimezoneApi, "/account/timezone")
|
|
||||||
api.add_resource(AccountPasswordApi, "/account/password")
|
|
||||||
api.add_resource(AccountIntegrateApi, "/account/integrates")
|
|
||||||
api.add_resource(AccountDeleteVerifyApi, "/account/delete/verify")
|
|
||||||
api.add_resource(AccountDeleteApi, "/account/delete")
|
|
||||||
api.add_resource(AccountDeleteUpdateFeedbackApi, "/account/delete/feedback")
|
|
||||||
api.add_resource(EducationVerifyApi, "/account/education/verify")
|
|
||||||
api.add_resource(EducationApi, "/account/education")
|
|
||||||
api.add_resource(EducationAutoCompleteApi, "/account/education/autocomplete")
|
|
||||||
# Change email
|
|
||||||
api.add_resource(ChangeEmailSendEmailApi, "/account/change-email")
|
|
||||||
api.add_resource(ChangeEmailCheckApi, "/account/change-email/validity")
|
|
||||||
api.add_resource(ChangeEmailResetApi, "/account/change-email/reset")
|
|
||||||
api.add_resource(CheckEmailUnique, "/account/change-email/check-email-unique")
|
|
||||||
# api.add_resource(AccountEmailApi, '/account/email')
|
|
||||||
# api.add_resource(AccountEmailVerifyApi, '/account/email-verify')
|
|
||||||
|
|||||||
@ -1,10 +1,10 @@
|
|||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, fields
|
from flask_restx import Resource, fields
|
||||||
|
|
||||||
from controllers.console import api, console_ns
|
from controllers.console import api, console_ns
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from core.model_runtime.utils.encoders import jsonable_encoder
|
from core.model_runtime.utils.encoders import jsonable_encoder
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from services.agent_service import AgentService
|
from services.agent_service import AgentService
|
||||||
|
|
||||||
|
|
||||||
@ -21,7 +21,9 @@ class AgentProviderListApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self):
|
def get(self):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
user = current_user
|
user = current_user
|
||||||
|
assert user.current_tenant_id is not None
|
||||||
|
|
||||||
user_id = user.id
|
user_id = user.id
|
||||||
tenant_id = user.current_tenant_id
|
tenant_id = user.current_tenant_id
|
||||||
@ -43,7 +45,9 @@ class AgentProviderApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self, provider_name: str):
|
def get(self, provider_name: str):
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
user = current_user
|
user = current_user
|
||||||
|
assert user.current_tenant_id is not None
|
||||||
user_id = user.id
|
user_id = user.id
|
||||||
tenant_id = user.current_tenant_id
|
tenant_id = user.current_tenant_id
|
||||||
return jsonable_encoder(AgentService.get_agent_provider(user_id, tenant_id, provider_name))
|
return jsonable_encoder(AgentService.get_agent_provider(user_id, tenant_id, provider_name))
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, fields, reqparse
|
from flask_restx import Resource, fields, reqparse
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
@ -6,10 +5,18 @@ from controllers.console import api, console_ns
|
|||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from core.model_runtime.utils.encoders import jsonable_encoder
|
from core.model_runtime.utils.encoders import jsonable_encoder
|
||||||
from core.plugin.impl.exc import PluginPermissionDeniedError
|
from core.plugin.impl.exc import PluginPermissionDeniedError
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
|
from models.account import Account
|
||||||
from services.plugin.endpoint_service import EndpointService
|
from services.plugin.endpoint_service import EndpointService
|
||||||
|
|
||||||
|
|
||||||
|
def _current_account_with_tenant() -> tuple[Account, str]:
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
tenant_id = current_user.current_tenant_id
|
||||||
|
assert tenant_id is not None
|
||||||
|
return current_user, tenant_id
|
||||||
|
|
||||||
|
|
||||||
@console_ns.route("/workspaces/current/endpoints/create")
|
@console_ns.route("/workspaces/current/endpoints/create")
|
||||||
class EndpointCreateApi(Resource):
|
class EndpointCreateApi(Resource):
|
||||||
@api.doc("create_endpoint")
|
@api.doc("create_endpoint")
|
||||||
@ -34,7 +41,7 @@ class EndpointCreateApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
if not user.is_admin_or_owner:
|
if not user.is_admin_or_owner:
|
||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
@ -51,7 +58,7 @@ class EndpointCreateApi(Resource):
|
|||||||
try:
|
try:
|
||||||
return {
|
return {
|
||||||
"success": EndpointService.create_endpoint(
|
"success": EndpointService.create_endpoint(
|
||||||
tenant_id=user.current_tenant_id,
|
tenant_id=tenant_id,
|
||||||
user_id=user.id,
|
user_id=user.id,
|
||||||
plugin_unique_identifier=plugin_unique_identifier,
|
plugin_unique_identifier=plugin_unique_identifier,
|
||||||
name=name,
|
name=name,
|
||||||
@ -80,7 +87,7 @@ class EndpointListApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self):
|
def get(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("page", type=int, required=True, location="args")
|
parser.add_argument("page", type=int, required=True, location="args")
|
||||||
@ -93,7 +100,7 @@ class EndpointListApi(Resource):
|
|||||||
return jsonable_encoder(
|
return jsonable_encoder(
|
||||||
{
|
{
|
||||||
"endpoints": EndpointService.list_endpoints(
|
"endpoints": EndpointService.list_endpoints(
|
||||||
tenant_id=user.current_tenant_id,
|
tenant_id=tenant_id,
|
||||||
user_id=user.id,
|
user_id=user.id,
|
||||||
page=page,
|
page=page,
|
||||||
page_size=page_size,
|
page_size=page_size,
|
||||||
@ -123,7 +130,7 @@ class EndpointListForSinglePluginApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def get(self):
|
def get(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("page", type=int, required=True, location="args")
|
parser.add_argument("page", type=int, required=True, location="args")
|
||||||
@ -138,7 +145,7 @@ class EndpointListForSinglePluginApi(Resource):
|
|||||||
return jsonable_encoder(
|
return jsonable_encoder(
|
||||||
{
|
{
|
||||||
"endpoints": EndpointService.list_endpoints_for_single_plugin(
|
"endpoints": EndpointService.list_endpoints_for_single_plugin(
|
||||||
tenant_id=user.current_tenant_id,
|
tenant_id=tenant_id,
|
||||||
user_id=user.id,
|
user_id=user.id,
|
||||||
plugin_id=plugin_id,
|
plugin_id=plugin_id,
|
||||||
page=page,
|
page=page,
|
||||||
@ -165,7 +172,7 @@ class EndpointDeleteApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("endpoint_id", type=str, required=True)
|
parser.add_argument("endpoint_id", type=str, required=True)
|
||||||
@ -177,9 +184,7 @@ class EndpointDeleteApi(Resource):
|
|||||||
endpoint_id = args["endpoint_id"]
|
endpoint_id = args["endpoint_id"]
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"success": EndpointService.delete_endpoint(
|
"success": EndpointService.delete_endpoint(tenant_id=tenant_id, user_id=user.id, endpoint_id=endpoint_id)
|
||||||
tenant_id=user.current_tenant_id, user_id=user.id, endpoint_id=endpoint_id
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@ -207,7 +212,7 @@ class EndpointUpdateApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("endpoint_id", type=str, required=True)
|
parser.add_argument("endpoint_id", type=str, required=True)
|
||||||
@ -224,7 +229,7 @@ class EndpointUpdateApi(Resource):
|
|||||||
|
|
||||||
return {
|
return {
|
||||||
"success": EndpointService.update_endpoint(
|
"success": EndpointService.update_endpoint(
|
||||||
tenant_id=user.current_tenant_id,
|
tenant_id=tenant_id,
|
||||||
user_id=user.id,
|
user_id=user.id,
|
||||||
endpoint_id=endpoint_id,
|
endpoint_id=endpoint_id,
|
||||||
name=name,
|
name=name,
|
||||||
@ -250,7 +255,7 @@ class EndpointEnableApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("endpoint_id", type=str, required=True)
|
parser.add_argument("endpoint_id", type=str, required=True)
|
||||||
@ -262,9 +267,7 @@ class EndpointEnableApi(Resource):
|
|||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"success": EndpointService.enable_endpoint(
|
"success": EndpointService.enable_endpoint(tenant_id=tenant_id, user_id=user.id, endpoint_id=endpoint_id)
|
||||||
tenant_id=user.current_tenant_id, user_id=user.id, endpoint_id=endpoint_id
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@ -285,7 +288,7 @@ class EndpointDisableApi(Resource):
|
|||||||
@login_required
|
@login_required
|
||||||
@account_initialization_required
|
@account_initialization_required
|
||||||
def post(self):
|
def post(self):
|
||||||
user = current_user
|
user, tenant_id = _current_account_with_tenant()
|
||||||
|
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
parser.add_argument("endpoint_id", type=str, required=True)
|
parser.add_argument("endpoint_id", type=str, required=True)
|
||||||
@ -297,7 +300,5 @@ class EndpointDisableApi(Resource):
|
|||||||
raise Forbidden()
|
raise Forbidden()
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"success": EndpointService.disable_endpoint(
|
"success": EndpointService.disable_endpoint(tenant_id=tenant_id, user_id=user.id, endpoint_id=endpoint_id)
|
||||||
tenant_id=user.current_tenant_id, user_id=user.id, endpoint_id=endpoint_id
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
from flask_restx import Resource, reqparse
|
from flask_restx import Resource, reqparse
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from core.model_runtime.entities.model_entities import ModelType
|
from core.model_runtime.entities.model_entities import ModelType
|
||||||
from core.model_runtime.errors.validate import CredentialsValidateFailedError
|
from core.model_runtime.errors.validate import CredentialsValidateFailedError
|
||||||
@ -10,6 +10,9 @@ from models.account import Account, TenantAccountRole
|
|||||||
from services.model_load_balancing_service import ModelLoadBalancingService
|
from services.model_load_balancing_service import ModelLoadBalancingService
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route(
|
||||||
|
"/workspaces/current/model-providers/<path:provider>/models/load-balancing-configs/credentials-validate"
|
||||||
|
)
|
||||||
class LoadBalancingCredentialsValidateApi(Resource):
|
class LoadBalancingCredentialsValidateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -61,6 +64,9 @@ class LoadBalancingCredentialsValidateApi(Resource):
|
|||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route(
|
||||||
|
"/workspaces/current/model-providers/<path:provider>/models/load-balancing-configs/<string:config_id>/credentials-validate"
|
||||||
|
)
|
||||||
class LoadBalancingConfigCredentialsValidateApi(Resource):
|
class LoadBalancingConfigCredentialsValidateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -111,15 +117,3 @@ class LoadBalancingConfigCredentialsValidateApi(Resource):
|
|||||||
response["error"] = error
|
response["error"] = error
|
||||||
|
|
||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
||||||
# Load Balancing Config
|
|
||||||
api.add_resource(
|
|
||||||
LoadBalancingCredentialsValidateApi,
|
|
||||||
"/workspaces/current/model-providers/<path:provider>/models/load-balancing-configs/credentials-validate",
|
|
||||||
)
|
|
||||||
|
|
||||||
api.add_resource(
|
|
||||||
LoadBalancingConfigCredentialsValidateApi,
|
|
||||||
"/workspaces/current/model-providers/<path:provider>/models/load-balancing-configs/<string:config_id>/credentials-validate",
|
|
||||||
)
|
|
||||||
|
|||||||
@ -1,12 +1,11 @@
|
|||||||
from urllib import parse
|
from urllib import parse
|
||||||
|
|
||||||
from flask import abort, request
|
from flask import abort, request
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, marshal_with, reqparse
|
from flask_restx import Resource, marshal_with, reqparse
|
||||||
|
|
||||||
import services
|
import services
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.auth.error import (
|
from controllers.console.auth.error import (
|
||||||
CannotTransferOwnerToSelfError,
|
CannotTransferOwnerToSelfError,
|
||||||
EmailCodeError,
|
EmailCodeError,
|
||||||
@ -26,13 +25,14 @@ from controllers.console.wraps import (
|
|||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from fields.member_fields import account_with_role_list_fields
|
from fields.member_fields import account_with_role_list_fields
|
||||||
from libs.helper import extract_remote_ip
|
from libs.helper import extract_remote_ip
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
from models.account import Account, TenantAccountRole
|
from models.account import Account, TenantAccountRole
|
||||||
from services.account_service import AccountService, RegisterService, TenantService
|
from services.account_service import AccountService, RegisterService, TenantService
|
||||||
from services.errors.account import AccountAlreadyInTenantError
|
from services.errors.account import AccountAlreadyInTenantError
|
||||||
from services.feature_service import FeatureService
|
from services.feature_service import FeatureService
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members")
|
||||||
class MemberListApi(Resource):
|
class MemberListApi(Resource):
|
||||||
"""List all members of current tenant."""
|
"""List all members of current tenant."""
|
||||||
|
|
||||||
@ -49,6 +49,7 @@ class MemberListApi(Resource):
|
|||||||
return {"result": "success", "accounts": members}, 200
|
return {"result": "success", "accounts": members}, 200
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members/invite-email")
|
||||||
class MemberInviteEmailApi(Resource):
|
class MemberInviteEmailApi(Resource):
|
||||||
"""Invite a new member by email."""
|
"""Invite a new member by email."""
|
||||||
|
|
||||||
@ -111,6 +112,7 @@ class MemberInviteEmailApi(Resource):
|
|||||||
}, 201
|
}, 201
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members/<uuid:member_id>")
|
||||||
class MemberCancelInviteApi(Resource):
|
class MemberCancelInviteApi(Resource):
|
||||||
"""Cancel an invitation by member id."""
|
"""Cancel an invitation by member id."""
|
||||||
|
|
||||||
@ -143,6 +145,7 @@ class MemberCancelInviteApi(Resource):
|
|||||||
}, 200
|
}, 200
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members/<uuid:member_id>/update-role")
|
||||||
class MemberUpdateRoleApi(Resource):
|
class MemberUpdateRoleApi(Resource):
|
||||||
"""Update member role."""
|
"""Update member role."""
|
||||||
|
|
||||||
@ -177,6 +180,7 @@ class MemberUpdateRoleApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/dataset-operators")
|
||||||
class DatasetOperatorMemberListApi(Resource):
|
class DatasetOperatorMemberListApi(Resource):
|
||||||
"""List all members of current tenant."""
|
"""List all members of current tenant."""
|
||||||
|
|
||||||
@ -193,6 +197,7 @@ class DatasetOperatorMemberListApi(Resource):
|
|||||||
return {"result": "success", "accounts": members}, 200
|
return {"result": "success", "accounts": members}, 200
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members/send-owner-transfer-confirm-email")
|
||||||
class SendOwnerTransferEmailApi(Resource):
|
class SendOwnerTransferEmailApi(Resource):
|
||||||
"""Send owner transfer email."""
|
"""Send owner transfer email."""
|
||||||
|
|
||||||
@ -233,6 +238,7 @@ class SendOwnerTransferEmailApi(Resource):
|
|||||||
return {"result": "success", "data": token}
|
return {"result": "success", "data": token}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members/owner-transfer-check")
|
||||||
class OwnerTransferCheckApi(Resource):
|
class OwnerTransferCheckApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -278,6 +284,7 @@ class OwnerTransferCheckApi(Resource):
|
|||||||
return {"is_valid": True, "email": token_data.get("email"), "token": new_token}
|
return {"is_valid": True, "email": token_data.get("email"), "token": new_token}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/members/<uuid:member_id>/owner-transfer")
|
||||||
class OwnerTransfer(Resource):
|
class OwnerTransfer(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -339,14 +346,3 @@ class OwnerTransfer(Resource):
|
|||||||
raise ValueError(str(e))
|
raise ValueError(str(e))
|
||||||
|
|
||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(MemberListApi, "/workspaces/current/members")
|
|
||||||
api.add_resource(MemberInviteEmailApi, "/workspaces/current/members/invite-email")
|
|
||||||
api.add_resource(MemberCancelInviteApi, "/workspaces/current/members/<uuid:member_id>")
|
|
||||||
api.add_resource(MemberUpdateRoleApi, "/workspaces/current/members/<uuid:member_id>/update-role")
|
|
||||||
api.add_resource(DatasetOperatorMemberListApi, "/workspaces/current/dataset-operators")
|
|
||||||
# owner transfer
|
|
||||||
api.add_resource(SendOwnerTransferEmailApi, "/workspaces/current/members/send-owner-transfer-confirm-email")
|
|
||||||
api.add_resource(OwnerTransferCheckApi, "/workspaces/current/members/owner-transfer-check")
|
|
||||||
api.add_resource(OwnerTransfer, "/workspaces/current/members/<uuid:member_id>/owner-transfer")
|
|
||||||
|
|||||||
@ -5,7 +5,7 @@ from flask_login import current_user
|
|||||||
from flask_restx import Resource, reqparse
|
from flask_restx import Resource, reqparse
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from core.model_runtime.entities.model_entities import ModelType
|
from core.model_runtime.entities.model_entities import ModelType
|
||||||
from core.model_runtime.errors.validate import CredentialsValidateFailedError
|
from core.model_runtime.errors.validate import CredentialsValidateFailedError
|
||||||
@ -17,6 +17,7 @@ from services.billing_service import BillingService
|
|||||||
from services.model_provider_service import ModelProviderService
|
from services.model_provider_service import ModelProviderService
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers")
|
||||||
class ModelProviderListApi(Resource):
|
class ModelProviderListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -45,6 +46,7 @@ class ModelProviderListApi(Resource):
|
|||||||
return jsonable_encoder({"data": provider_list})
|
return jsonable_encoder({"data": provider_list})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/credentials")
|
||||||
class ModelProviderCredentialApi(Resource):
|
class ModelProviderCredentialApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -151,6 +153,7 @@ class ModelProviderCredentialApi(Resource):
|
|||||||
return {"result": "success"}, 204
|
return {"result": "success"}, 204
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/credentials/switch")
|
||||||
class ModelProviderCredentialSwitchApi(Resource):
|
class ModelProviderCredentialSwitchApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -175,6 +178,7 @@ class ModelProviderCredentialSwitchApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/credentials/validate")
|
||||||
class ModelProviderValidateApi(Resource):
|
class ModelProviderValidateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -211,6 +215,7 @@ class ModelProviderValidateApi(Resource):
|
|||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/<string:tenant_id>/model-providers/<path:provider>/<string:icon_type>/<string:lang>")
|
||||||
class ModelProviderIconApi(Resource):
|
class ModelProviderIconApi(Resource):
|
||||||
"""
|
"""
|
||||||
Get model provider icon
|
Get model provider icon
|
||||||
@ -229,6 +234,7 @@ class ModelProviderIconApi(Resource):
|
|||||||
return send_file(io.BytesIO(icon), mimetype=mimetype)
|
return send_file(io.BytesIO(icon), mimetype=mimetype)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/preferred-provider-type")
|
||||||
class PreferredProviderTypeUpdateApi(Resource):
|
class PreferredProviderTypeUpdateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -262,6 +268,7 @@ class PreferredProviderTypeUpdateApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/checkout-url")
|
||||||
class ModelProviderPaymentCheckoutUrlApi(Resource):
|
class ModelProviderPaymentCheckoutUrlApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -281,21 +288,3 @@ class ModelProviderPaymentCheckoutUrlApi(Resource):
|
|||||||
prefilled_email=current_user.email,
|
prefilled_email=current_user.email,
|
||||||
)
|
)
|
||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(ModelProviderListApi, "/workspaces/current/model-providers")
|
|
||||||
|
|
||||||
api.add_resource(ModelProviderCredentialApi, "/workspaces/current/model-providers/<path:provider>/credentials")
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderCredentialSwitchApi, "/workspaces/current/model-providers/<path:provider>/credentials/switch"
|
|
||||||
)
|
|
||||||
api.add_resource(ModelProviderValidateApi, "/workspaces/current/model-providers/<path:provider>/credentials/validate")
|
|
||||||
|
|
||||||
api.add_resource(
|
|
||||||
PreferredProviderTypeUpdateApi, "/workspaces/current/model-providers/<path:provider>/preferred-provider-type"
|
|
||||||
)
|
|
||||||
api.add_resource(ModelProviderPaymentCheckoutUrlApi, "/workspaces/current/model-providers/<path:provider>/checkout-url")
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderIconApi,
|
|
||||||
"/workspaces/<string:tenant_id>/model-providers/<path:provider>/<string:icon_type>/<string:lang>",
|
|
||||||
)
|
|
||||||
|
|||||||
@ -4,7 +4,7 @@ from flask_login import current_user
|
|||||||
from flask_restx import Resource, reqparse
|
from flask_restx import Resource, reqparse
|
||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from core.model_runtime.entities.model_entities import ModelType
|
from core.model_runtime.entities.model_entities import ModelType
|
||||||
from core.model_runtime.errors.validate import CredentialsValidateFailedError
|
from core.model_runtime.errors.validate import CredentialsValidateFailedError
|
||||||
@ -17,6 +17,7 @@ from services.model_provider_service import ModelProviderService
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/default-model")
|
||||||
class DefaultModelApi(Resource):
|
class DefaultModelApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -85,6 +86,7 @@ class DefaultModelApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/models")
|
||||||
class ModelProviderModelApi(Resource):
|
class ModelProviderModelApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -187,6 +189,7 @@ class ModelProviderModelApi(Resource):
|
|||||||
return {"result": "success"}, 204
|
return {"result": "success"}, 204
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/models/credentials")
|
||||||
class ModelProviderModelCredentialApi(Resource):
|
class ModelProviderModelCredentialApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -364,6 +367,7 @@ class ModelProviderModelCredentialApi(Resource):
|
|||||||
return {"result": "success"}, 204
|
return {"result": "success"}, 204
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/models/credentials/switch")
|
||||||
class ModelProviderModelCredentialSwitchApi(Resource):
|
class ModelProviderModelCredentialSwitchApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -395,6 +399,9 @@ class ModelProviderModelCredentialSwitchApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route(
|
||||||
|
"/workspaces/current/model-providers/<path:provider>/models/enable", endpoint="model-provider-model-enable"
|
||||||
|
)
|
||||||
class ModelProviderModelEnableApi(Resource):
|
class ModelProviderModelEnableApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -422,6 +429,9 @@ class ModelProviderModelEnableApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route(
|
||||||
|
"/workspaces/current/model-providers/<path:provider>/models/disable", endpoint="model-provider-model-disable"
|
||||||
|
)
|
||||||
class ModelProviderModelDisableApi(Resource):
|
class ModelProviderModelDisableApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -449,6 +459,7 @@ class ModelProviderModelDisableApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/models/credentials/validate")
|
||||||
class ModelProviderModelValidateApi(Resource):
|
class ModelProviderModelValidateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -494,6 +505,7 @@ class ModelProviderModelValidateApi(Resource):
|
|||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/model-providers/<path:provider>/models/parameter-rules")
|
||||||
class ModelProviderModelParameterRuleApi(Resource):
|
class ModelProviderModelParameterRuleApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -513,6 +525,7 @@ class ModelProviderModelParameterRuleApi(Resource):
|
|||||||
return jsonable_encoder({"data": parameter_rules})
|
return jsonable_encoder({"data": parameter_rules})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/models/model-types/<string:model_type>")
|
||||||
class ModelProviderAvailableModelApi(Resource):
|
class ModelProviderAvailableModelApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -524,32 +537,3 @@ class ModelProviderAvailableModelApi(Resource):
|
|||||||
models = model_provider_service.get_models_by_model_type(tenant_id=tenant_id, model_type=model_type)
|
models = model_provider_service.get_models_by_model_type(tenant_id=tenant_id, model_type=model_type)
|
||||||
|
|
||||||
return jsonable_encoder({"data": models})
|
return jsonable_encoder({"data": models})
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(ModelProviderModelApi, "/workspaces/current/model-providers/<path:provider>/models")
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderModelEnableApi,
|
|
||||||
"/workspaces/current/model-providers/<path:provider>/models/enable",
|
|
||||||
endpoint="model-provider-model-enable",
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderModelDisableApi,
|
|
||||||
"/workspaces/current/model-providers/<path:provider>/models/disable",
|
|
||||||
endpoint="model-provider-model-disable",
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderModelCredentialApi, "/workspaces/current/model-providers/<path:provider>/models/credentials"
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderModelCredentialSwitchApi,
|
|
||||||
"/workspaces/current/model-providers/<path:provider>/models/credentials/switch",
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderModelValidateApi, "/workspaces/current/model-providers/<path:provider>/models/credentials/validate"
|
|
||||||
)
|
|
||||||
|
|
||||||
api.add_resource(
|
|
||||||
ModelProviderModelParameterRuleApi, "/workspaces/current/model-providers/<path:provider>/models/parameter-rules"
|
|
||||||
)
|
|
||||||
api.add_resource(ModelProviderAvailableModelApi, "/workspaces/current/models/model-types/<string:model_type>")
|
|
||||||
api.add_resource(DefaultModelApi, "/workspaces/current/default-model")
|
|
||||||
|
|||||||
@ -6,7 +6,7 @@ from flask_restx import Resource, reqparse
|
|||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.workspace import plugin_permission_required
|
from controllers.console.workspace import plugin_permission_required
|
||||||
from controllers.console.wraps import account_initialization_required, setup_required
|
from controllers.console.wraps import account_initialization_required, setup_required
|
||||||
from core.model_runtime.utils.encoders import jsonable_encoder
|
from core.model_runtime.utils.encoders import jsonable_encoder
|
||||||
@ -19,6 +19,7 @@ from services.plugin.plugin_permission_service import PluginPermissionService
|
|||||||
from services.plugin.plugin_service import PluginService
|
from services.plugin.plugin_service import PluginService
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/debugging-key")
|
||||||
class PluginDebuggingKeyApi(Resource):
|
class PluginDebuggingKeyApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -37,6 +38,7 @@ class PluginDebuggingKeyApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/list")
|
||||||
class PluginListApi(Resource):
|
class PluginListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -55,6 +57,7 @@ class PluginListApi(Resource):
|
|||||||
return jsonable_encoder({"plugins": plugins_with_total.list, "total": plugins_with_total.total})
|
return jsonable_encoder({"plugins": plugins_with_total.list, "total": plugins_with_total.total})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/list/latest-versions")
|
||||||
class PluginListLatestVersionsApi(Resource):
|
class PluginListLatestVersionsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -72,6 +75,7 @@ class PluginListLatestVersionsApi(Resource):
|
|||||||
return jsonable_encoder({"versions": versions})
|
return jsonable_encoder({"versions": versions})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/list/installations/ids")
|
||||||
class PluginListInstallationsFromIdsApi(Resource):
|
class PluginListInstallationsFromIdsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -91,6 +95,7 @@ class PluginListInstallationsFromIdsApi(Resource):
|
|||||||
return jsonable_encoder({"plugins": plugins})
|
return jsonable_encoder({"plugins": plugins})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/icon")
|
||||||
class PluginIconApi(Resource):
|
class PluginIconApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
def get(self):
|
def get(self):
|
||||||
@ -108,6 +113,7 @@ class PluginIconApi(Resource):
|
|||||||
return send_file(io.BytesIO(icon_bytes), mimetype=mimetype, max_age=icon_cache_max_age)
|
return send_file(io.BytesIO(icon_bytes), mimetype=mimetype, max_age=icon_cache_max_age)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/upload/pkg")
|
||||||
class PluginUploadFromPkgApi(Resource):
|
class PluginUploadFromPkgApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -131,6 +137,7 @@ class PluginUploadFromPkgApi(Resource):
|
|||||||
return jsonable_encoder(response)
|
return jsonable_encoder(response)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/upload/github")
|
||||||
class PluginUploadFromGithubApi(Resource):
|
class PluginUploadFromGithubApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -153,6 +160,7 @@ class PluginUploadFromGithubApi(Resource):
|
|||||||
return jsonable_encoder(response)
|
return jsonable_encoder(response)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/upload/bundle")
|
||||||
class PluginUploadFromBundleApi(Resource):
|
class PluginUploadFromBundleApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -176,6 +184,7 @@ class PluginUploadFromBundleApi(Resource):
|
|||||||
return jsonable_encoder(response)
|
return jsonable_encoder(response)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/install/pkg")
|
||||||
class PluginInstallFromPkgApi(Resource):
|
class PluginInstallFromPkgApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -201,6 +210,7 @@ class PluginInstallFromPkgApi(Resource):
|
|||||||
return jsonable_encoder(response)
|
return jsonable_encoder(response)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/install/github")
|
||||||
class PluginInstallFromGithubApi(Resource):
|
class PluginInstallFromGithubApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -230,6 +240,7 @@ class PluginInstallFromGithubApi(Resource):
|
|||||||
return jsonable_encoder(response)
|
return jsonable_encoder(response)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/install/marketplace")
|
||||||
class PluginInstallFromMarketplaceApi(Resource):
|
class PluginInstallFromMarketplaceApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -255,6 +266,7 @@ class PluginInstallFromMarketplaceApi(Resource):
|
|||||||
return jsonable_encoder(response)
|
return jsonable_encoder(response)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/marketplace/pkg")
|
||||||
class PluginFetchMarketplacePkgApi(Resource):
|
class PluginFetchMarketplacePkgApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -280,6 +292,7 @@ class PluginFetchMarketplacePkgApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/fetch-manifest")
|
||||||
class PluginFetchManifestApi(Resource):
|
class PluginFetchManifestApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -304,6 +317,7 @@ class PluginFetchManifestApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/tasks")
|
||||||
class PluginFetchInstallTasksApi(Resource):
|
class PluginFetchInstallTasksApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -325,6 +339,7 @@ class PluginFetchInstallTasksApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/tasks/<task_id>")
|
||||||
class PluginFetchInstallTaskApi(Resource):
|
class PluginFetchInstallTaskApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -339,6 +354,7 @@ class PluginFetchInstallTaskApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/tasks/<task_id>/delete")
|
||||||
class PluginDeleteInstallTaskApi(Resource):
|
class PluginDeleteInstallTaskApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -353,6 +369,7 @@ class PluginDeleteInstallTaskApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/tasks/delete_all")
|
||||||
class PluginDeleteAllInstallTaskItemsApi(Resource):
|
class PluginDeleteAllInstallTaskItemsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -367,6 +384,7 @@ class PluginDeleteAllInstallTaskItemsApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/tasks/<task_id>/delete/<path:identifier>")
|
||||||
class PluginDeleteInstallTaskItemApi(Resource):
|
class PluginDeleteInstallTaskItemApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -381,6 +399,7 @@ class PluginDeleteInstallTaskItemApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/upgrade/marketplace")
|
||||||
class PluginUpgradeFromMarketplaceApi(Resource):
|
class PluginUpgradeFromMarketplaceApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -404,6 +423,7 @@ class PluginUpgradeFromMarketplaceApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/upgrade/github")
|
||||||
class PluginUpgradeFromGithubApi(Resource):
|
class PluginUpgradeFromGithubApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -435,6 +455,7 @@ class PluginUpgradeFromGithubApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/uninstall")
|
||||||
class PluginUninstallApi(Resource):
|
class PluginUninstallApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -453,6 +474,7 @@ class PluginUninstallApi(Resource):
|
|||||||
raise ValueError(e)
|
raise ValueError(e)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/permission/change")
|
||||||
class PluginChangePermissionApi(Resource):
|
class PluginChangePermissionApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -475,6 +497,7 @@ class PluginChangePermissionApi(Resource):
|
|||||||
return {"success": PluginPermissionService.change_permission(tenant_id, install_permission, debug_permission)}
|
return {"success": PluginPermissionService.change_permission(tenant_id, install_permission, debug_permission)}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/permission/fetch")
|
||||||
class PluginFetchPermissionApi(Resource):
|
class PluginFetchPermissionApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -499,6 +522,7 @@ class PluginFetchPermissionApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/parameters/dynamic-options")
|
||||||
class PluginFetchDynamicSelectOptionsApi(Resource):
|
class PluginFetchDynamicSelectOptionsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -535,6 +559,7 @@ class PluginFetchDynamicSelectOptionsApi(Resource):
|
|||||||
return jsonable_encoder({"options": options})
|
return jsonable_encoder({"options": options})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/preferences/change")
|
||||||
class PluginChangePreferencesApi(Resource):
|
class PluginChangePreferencesApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -590,6 +615,7 @@ class PluginChangePreferencesApi(Resource):
|
|||||||
return jsonable_encoder({"success": True})
|
return jsonable_encoder({"success": True})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/preferences/fetch")
|
||||||
class PluginFetchPreferencesApi(Resource):
|
class PluginFetchPreferencesApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -628,6 +654,7 @@ class PluginFetchPreferencesApi(Resource):
|
|||||||
return jsonable_encoder({"permission": permission_dict, "auto_upgrade": auto_upgrade_dict})
|
return jsonable_encoder({"permission": permission_dict, "auto_upgrade": auto_upgrade_dict})
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/plugin/preferences/autoupgrade/exclude")
|
||||||
class PluginAutoUpgradeExcludePluginApi(Resource):
|
class PluginAutoUpgradeExcludePluginApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -641,35 +668,3 @@ class PluginAutoUpgradeExcludePluginApi(Resource):
|
|||||||
args = req.parse_args()
|
args = req.parse_args()
|
||||||
|
|
||||||
return jsonable_encoder({"success": PluginAutoUpgradeService.exclude_plugin(tenant_id, args["plugin_id"])})
|
return jsonable_encoder({"success": PluginAutoUpgradeService.exclude_plugin(tenant_id, args["plugin_id"])})
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(PluginDebuggingKeyApi, "/workspaces/current/plugin/debugging-key")
|
|
||||||
api.add_resource(PluginListApi, "/workspaces/current/plugin/list")
|
|
||||||
api.add_resource(PluginListLatestVersionsApi, "/workspaces/current/plugin/list/latest-versions")
|
|
||||||
api.add_resource(PluginListInstallationsFromIdsApi, "/workspaces/current/plugin/list/installations/ids")
|
|
||||||
api.add_resource(PluginIconApi, "/workspaces/current/plugin/icon")
|
|
||||||
api.add_resource(PluginUploadFromPkgApi, "/workspaces/current/plugin/upload/pkg")
|
|
||||||
api.add_resource(PluginUploadFromGithubApi, "/workspaces/current/plugin/upload/github")
|
|
||||||
api.add_resource(PluginUploadFromBundleApi, "/workspaces/current/plugin/upload/bundle")
|
|
||||||
api.add_resource(PluginInstallFromPkgApi, "/workspaces/current/plugin/install/pkg")
|
|
||||||
api.add_resource(PluginInstallFromGithubApi, "/workspaces/current/plugin/install/github")
|
|
||||||
api.add_resource(PluginUpgradeFromMarketplaceApi, "/workspaces/current/plugin/upgrade/marketplace")
|
|
||||||
api.add_resource(PluginUpgradeFromGithubApi, "/workspaces/current/plugin/upgrade/github")
|
|
||||||
api.add_resource(PluginInstallFromMarketplaceApi, "/workspaces/current/plugin/install/marketplace")
|
|
||||||
api.add_resource(PluginFetchManifestApi, "/workspaces/current/plugin/fetch-manifest")
|
|
||||||
api.add_resource(PluginFetchInstallTasksApi, "/workspaces/current/plugin/tasks")
|
|
||||||
api.add_resource(PluginFetchInstallTaskApi, "/workspaces/current/plugin/tasks/<task_id>")
|
|
||||||
api.add_resource(PluginDeleteInstallTaskApi, "/workspaces/current/plugin/tasks/<task_id>/delete")
|
|
||||||
api.add_resource(PluginDeleteAllInstallTaskItemsApi, "/workspaces/current/plugin/tasks/delete_all")
|
|
||||||
api.add_resource(PluginDeleteInstallTaskItemApi, "/workspaces/current/plugin/tasks/<task_id>/delete/<path:identifier>")
|
|
||||||
api.add_resource(PluginUninstallApi, "/workspaces/current/plugin/uninstall")
|
|
||||||
api.add_resource(PluginFetchMarketplacePkgApi, "/workspaces/current/plugin/marketplace/pkg")
|
|
||||||
|
|
||||||
api.add_resource(PluginChangePermissionApi, "/workspaces/current/plugin/permission/change")
|
|
||||||
api.add_resource(PluginFetchPermissionApi, "/workspaces/current/plugin/permission/fetch")
|
|
||||||
|
|
||||||
api.add_resource(PluginFetchDynamicSelectOptionsApi, "/workspaces/current/plugin/parameters/dynamic-options")
|
|
||||||
|
|
||||||
api.add_resource(PluginFetchPreferencesApi, "/workspaces/current/plugin/preferences/fetch")
|
|
||||||
api.add_resource(PluginChangePreferencesApi, "/workspaces/current/plugin/preferences/change")
|
|
||||||
api.add_resource(PluginAutoUpgradeExcludePluginApi, "/workspaces/current/plugin/preferences/autoupgrade/exclude")
|
|
||||||
|
|||||||
@ -11,7 +11,7 @@ from sqlalchemy.orm import Session
|
|||||||
from werkzeug.exceptions import Forbidden
|
from werkzeug.exceptions import Forbidden
|
||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.wraps import (
|
from controllers.console.wraps import (
|
||||||
account_initialization_required,
|
account_initialization_required,
|
||||||
enterprise_license_required,
|
enterprise_license_required,
|
||||||
@ -48,6 +48,7 @@ def is_valid_url(url: str) -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-providers")
|
||||||
class ToolProviderListApi(Resource):
|
class ToolProviderListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -72,6 +73,7 @@ class ToolProviderListApi(Resource):
|
|||||||
return ToolCommonService.list_tool_providers(user_id, tenant_id, args.get("type", None))
|
return ToolCommonService.list_tool_providers(user_id, tenant_id, args.get("type", None))
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/tools")
|
||||||
class ToolBuiltinProviderListToolsApi(Resource):
|
class ToolBuiltinProviderListToolsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -89,6 +91,7 @@ class ToolBuiltinProviderListToolsApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/info")
|
||||||
class ToolBuiltinProviderInfoApi(Resource):
|
class ToolBuiltinProviderInfoApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -101,6 +104,7 @@ class ToolBuiltinProviderInfoApi(Resource):
|
|||||||
return jsonable_encoder(BuiltinToolManageService.get_builtin_tool_provider_info(tenant_id, provider))
|
return jsonable_encoder(BuiltinToolManageService.get_builtin_tool_provider_info(tenant_id, provider))
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/delete")
|
||||||
class ToolBuiltinProviderDeleteApi(Resource):
|
class ToolBuiltinProviderDeleteApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -122,6 +126,7 @@ class ToolBuiltinProviderDeleteApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/add")
|
||||||
class ToolBuiltinProviderAddApi(Resource):
|
class ToolBuiltinProviderAddApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -151,6 +156,7 @@ class ToolBuiltinProviderAddApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/update")
|
||||||
class ToolBuiltinProviderUpdateApi(Resource):
|
class ToolBuiltinProviderUpdateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -182,6 +188,7 @@ class ToolBuiltinProviderUpdateApi(Resource):
|
|||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/credentials")
|
||||||
class ToolBuiltinProviderGetCredentialsApi(Resource):
|
class ToolBuiltinProviderGetCredentialsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -197,6 +204,7 @@ class ToolBuiltinProviderGetCredentialsApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/icon")
|
||||||
class ToolBuiltinProviderIconApi(Resource):
|
class ToolBuiltinProviderIconApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
def get(self, provider):
|
def get(self, provider):
|
||||||
@ -205,6 +213,7 @@ class ToolBuiltinProviderIconApi(Resource):
|
|||||||
return send_file(io.BytesIO(icon_bytes), mimetype=mimetype, max_age=icon_cache_max_age)
|
return send_file(io.BytesIO(icon_bytes), mimetype=mimetype, max_age=icon_cache_max_age)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/add")
|
||||||
class ToolApiProviderAddApi(Resource):
|
class ToolApiProviderAddApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -244,6 +253,7 @@ class ToolApiProviderAddApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/remote")
|
||||||
class ToolApiProviderGetRemoteSchemaApi(Resource):
|
class ToolApiProviderGetRemoteSchemaApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -267,6 +277,7 @@ class ToolApiProviderGetRemoteSchemaApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/tools")
|
||||||
class ToolApiProviderListToolsApi(Resource):
|
class ToolApiProviderListToolsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -292,6 +303,7 @@ class ToolApiProviderListToolsApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/update")
|
||||||
class ToolApiProviderUpdateApi(Resource):
|
class ToolApiProviderUpdateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -333,6 +345,7 @@ class ToolApiProviderUpdateApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/delete")
|
||||||
class ToolApiProviderDeleteApi(Resource):
|
class ToolApiProviderDeleteApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -359,6 +372,7 @@ class ToolApiProviderDeleteApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/get")
|
||||||
class ToolApiProviderGetApi(Resource):
|
class ToolApiProviderGetApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -382,6 +396,7 @@ class ToolApiProviderGetApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/credential/schema/<path:credential_type>")
|
||||||
class ToolBuiltinProviderCredentialsSchemaApi(Resource):
|
class ToolBuiltinProviderCredentialsSchemaApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -397,6 +412,7 @@ class ToolBuiltinProviderCredentialsSchemaApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/schema")
|
||||||
class ToolApiProviderSchemaApi(Resource):
|
class ToolApiProviderSchemaApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -413,6 +429,7 @@ class ToolApiProviderSchemaApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/api/test/pre")
|
||||||
class ToolApiProviderPreviousTestApi(Resource):
|
class ToolApiProviderPreviousTestApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -440,6 +457,7 @@ class ToolApiProviderPreviousTestApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/workflow/create")
|
||||||
class ToolWorkflowProviderCreateApi(Resource):
|
class ToolWorkflowProviderCreateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -479,6 +497,7 @@ class ToolWorkflowProviderCreateApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/workflow/update")
|
||||||
class ToolWorkflowProviderUpdateApi(Resource):
|
class ToolWorkflowProviderUpdateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -521,6 +540,7 @@ class ToolWorkflowProviderUpdateApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/workflow/delete")
|
||||||
class ToolWorkflowProviderDeleteApi(Resource):
|
class ToolWorkflowProviderDeleteApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -546,6 +566,7 @@ class ToolWorkflowProviderDeleteApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/workflow/get")
|
||||||
class ToolWorkflowProviderGetApi(Resource):
|
class ToolWorkflowProviderGetApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -580,6 +601,7 @@ class ToolWorkflowProviderGetApi(Resource):
|
|||||||
return jsonable_encoder(tool)
|
return jsonable_encoder(tool)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/workflow/tools")
|
||||||
class ToolWorkflowProviderListToolApi(Resource):
|
class ToolWorkflowProviderListToolApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -604,6 +626,7 @@ class ToolWorkflowProviderListToolApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tools/builtin")
|
||||||
class ToolBuiltinListApi(Resource):
|
class ToolBuiltinListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -625,6 +648,7 @@ class ToolBuiltinListApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tools/api")
|
||||||
class ToolApiListApi(Resource):
|
class ToolApiListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -643,6 +667,7 @@ class ToolApiListApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tools/workflow")
|
||||||
class ToolWorkflowListApi(Resource):
|
class ToolWorkflowListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -664,6 +689,7 @@ class ToolWorkflowListApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-labels")
|
||||||
class ToolLabelsApi(Resource):
|
class ToolLabelsApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -673,6 +699,7 @@ class ToolLabelsApi(Resource):
|
|||||||
return jsonable_encoder(ToolLabelsService.list_tool_labels())
|
return jsonable_encoder(ToolLabelsService.list_tool_labels())
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/oauth/plugin/<path:provider>/tool/authorization-url")
|
||||||
class ToolPluginOAuthApi(Resource):
|
class ToolPluginOAuthApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -717,6 +744,7 @@ class ToolPluginOAuthApi(Resource):
|
|||||||
return response
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/oauth/plugin/<path:provider>/tool/callback")
|
||||||
class ToolOAuthCallback(Resource):
|
class ToolOAuthCallback(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
def get(self, provider):
|
def get(self, provider):
|
||||||
@ -767,6 +795,7 @@ class ToolOAuthCallback(Resource):
|
|||||||
return redirect(f"{dify_config.CONSOLE_WEB_URL}/oauth-callback")
|
return redirect(f"{dify_config.CONSOLE_WEB_URL}/oauth-callback")
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/default-credential")
|
||||||
class ToolBuiltinProviderSetDefaultApi(Resource):
|
class ToolBuiltinProviderSetDefaultApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -780,6 +809,7 @@ class ToolBuiltinProviderSetDefaultApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/oauth/custom-client")
|
||||||
class ToolOAuthCustomClient(Resource):
|
class ToolOAuthCustomClient(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -823,6 +853,7 @@ class ToolOAuthCustomClient(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/oauth/client-schema")
|
||||||
class ToolBuiltinProviderGetOauthClientSchemaApi(Resource):
|
class ToolBuiltinProviderGetOauthClientSchemaApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -835,6 +866,7 @@ class ToolBuiltinProviderGetOauthClientSchemaApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/builtin/<path:provider>/credential/info")
|
||||||
class ToolBuiltinProviderGetCredentialInfoApi(Resource):
|
class ToolBuiltinProviderGetCredentialInfoApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -850,6 +882,7 @@ class ToolBuiltinProviderGetCredentialInfoApi(Resource):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/mcp")
|
||||||
class ToolProviderMCPApi(Resource):
|
class ToolProviderMCPApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -960,6 +993,7 @@ class ToolProviderMCPApi(Resource):
|
|||||||
return {"result": "success"}
|
return {"result": "success"}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/mcp/auth")
|
||||||
class ToolMCPAuthApi(Resource):
|
class ToolMCPAuthApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -1013,6 +1047,7 @@ class ToolMCPAuthApi(Resource):
|
|||||||
raise ValueError(f"Failed to connect to MCP server: {e}") from e
|
raise ValueError(f"Failed to connect to MCP server: {e}") from e
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/mcp/tools/<path:provider_id>")
|
||||||
class ToolMCPDetailApi(Resource):
|
class ToolMCPDetailApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -1025,6 +1060,7 @@ class ToolMCPDetailApi(Resource):
|
|||||||
return jsonable_encoder(ToolTransformService.mcp_provider_to_user_provider(provider, for_list=True))
|
return jsonable_encoder(ToolTransformService.mcp_provider_to_user_provider(provider, for_list=True))
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tools/mcp")
|
||||||
class ToolMCPListAllApi(Resource):
|
class ToolMCPListAllApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -1040,6 +1076,7 @@ class ToolMCPListAllApi(Resource):
|
|||||||
return [tool.to_dict() for tool in tools]
|
return [tool.to_dict() for tool in tools]
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current/tool-provider/mcp/update/<path:provider_id>")
|
||||||
class ToolMCPUpdateApi(Resource):
|
class ToolMCPUpdateApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -1055,6 +1092,7 @@ class ToolMCPUpdateApi(Resource):
|
|||||||
return jsonable_encoder(tools)
|
return jsonable_encoder(tools)
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/mcp/oauth/callback")
|
||||||
class ToolMCPCallbackApi(Resource):
|
class ToolMCPCallbackApi(Resource):
|
||||||
def get(self):
|
def get(self):
|
||||||
parser = reqparse.RequestParser()
|
parser = reqparse.RequestParser()
|
||||||
@ -1071,67 +1109,3 @@ class ToolMCPCallbackApi(Resource):
|
|||||||
session.commit()
|
session.commit()
|
||||||
|
|
||||||
return redirect(f"{dify_config.CONSOLE_WEB_URL}/oauth-callback")
|
return redirect(f"{dify_config.CONSOLE_WEB_URL}/oauth-callback")
|
||||||
|
|
||||||
|
|
||||||
# tool provider
|
|
||||||
api.add_resource(ToolProviderListApi, "/workspaces/current/tool-providers")
|
|
||||||
|
|
||||||
# tool oauth
|
|
||||||
api.add_resource(ToolPluginOAuthApi, "/oauth/plugin/<path:provider>/tool/authorization-url")
|
|
||||||
api.add_resource(ToolOAuthCallback, "/oauth/plugin/<path:provider>/tool/callback")
|
|
||||||
api.add_resource(ToolOAuthCustomClient, "/workspaces/current/tool-provider/builtin/<path:provider>/oauth/custom-client")
|
|
||||||
|
|
||||||
# builtin tool provider
|
|
||||||
api.add_resource(ToolBuiltinProviderListToolsApi, "/workspaces/current/tool-provider/builtin/<path:provider>/tools")
|
|
||||||
api.add_resource(ToolBuiltinProviderInfoApi, "/workspaces/current/tool-provider/builtin/<path:provider>/info")
|
|
||||||
api.add_resource(ToolBuiltinProviderAddApi, "/workspaces/current/tool-provider/builtin/<path:provider>/add")
|
|
||||||
api.add_resource(ToolBuiltinProviderDeleteApi, "/workspaces/current/tool-provider/builtin/<path:provider>/delete")
|
|
||||||
api.add_resource(ToolBuiltinProviderUpdateApi, "/workspaces/current/tool-provider/builtin/<path:provider>/update")
|
|
||||||
api.add_resource(
|
|
||||||
ToolBuiltinProviderSetDefaultApi, "/workspaces/current/tool-provider/builtin/<path:provider>/default-credential"
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ToolBuiltinProviderGetCredentialInfoApi, "/workspaces/current/tool-provider/builtin/<path:provider>/credential/info"
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ToolBuiltinProviderGetCredentialsApi, "/workspaces/current/tool-provider/builtin/<path:provider>/credentials"
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ToolBuiltinProviderCredentialsSchemaApi,
|
|
||||||
"/workspaces/current/tool-provider/builtin/<path:provider>/credential/schema/<path:credential_type>",
|
|
||||||
)
|
|
||||||
api.add_resource(
|
|
||||||
ToolBuiltinProviderGetOauthClientSchemaApi,
|
|
||||||
"/workspaces/current/tool-provider/builtin/<path:provider>/oauth/client-schema",
|
|
||||||
)
|
|
||||||
api.add_resource(ToolBuiltinProviderIconApi, "/workspaces/current/tool-provider/builtin/<path:provider>/icon")
|
|
||||||
|
|
||||||
# api tool provider
|
|
||||||
api.add_resource(ToolApiProviderAddApi, "/workspaces/current/tool-provider/api/add")
|
|
||||||
api.add_resource(ToolApiProviderGetRemoteSchemaApi, "/workspaces/current/tool-provider/api/remote")
|
|
||||||
api.add_resource(ToolApiProviderListToolsApi, "/workspaces/current/tool-provider/api/tools")
|
|
||||||
api.add_resource(ToolApiProviderUpdateApi, "/workspaces/current/tool-provider/api/update")
|
|
||||||
api.add_resource(ToolApiProviderDeleteApi, "/workspaces/current/tool-provider/api/delete")
|
|
||||||
api.add_resource(ToolApiProviderGetApi, "/workspaces/current/tool-provider/api/get")
|
|
||||||
api.add_resource(ToolApiProviderSchemaApi, "/workspaces/current/tool-provider/api/schema")
|
|
||||||
api.add_resource(ToolApiProviderPreviousTestApi, "/workspaces/current/tool-provider/api/test/pre")
|
|
||||||
|
|
||||||
# workflow tool provider
|
|
||||||
api.add_resource(ToolWorkflowProviderCreateApi, "/workspaces/current/tool-provider/workflow/create")
|
|
||||||
api.add_resource(ToolWorkflowProviderUpdateApi, "/workspaces/current/tool-provider/workflow/update")
|
|
||||||
api.add_resource(ToolWorkflowProviderDeleteApi, "/workspaces/current/tool-provider/workflow/delete")
|
|
||||||
api.add_resource(ToolWorkflowProviderGetApi, "/workspaces/current/tool-provider/workflow/get")
|
|
||||||
api.add_resource(ToolWorkflowProviderListToolApi, "/workspaces/current/tool-provider/workflow/tools")
|
|
||||||
|
|
||||||
# mcp tool provider
|
|
||||||
api.add_resource(ToolMCPDetailApi, "/workspaces/current/tool-provider/mcp/tools/<path:provider_id>")
|
|
||||||
api.add_resource(ToolProviderMCPApi, "/workspaces/current/tool-provider/mcp")
|
|
||||||
api.add_resource(ToolMCPUpdateApi, "/workspaces/current/tool-provider/mcp/update/<path:provider_id>")
|
|
||||||
api.add_resource(ToolMCPAuthApi, "/workspaces/current/tool-provider/mcp/auth")
|
|
||||||
api.add_resource(ToolMCPCallbackApi, "/mcp/oauth/callback")
|
|
||||||
|
|
||||||
api.add_resource(ToolBuiltinListApi, "/workspaces/current/tools/builtin")
|
|
||||||
api.add_resource(ToolApiListApi, "/workspaces/current/tools/api")
|
|
||||||
api.add_resource(ToolMCPListAllApi, "/workspaces/current/tools/mcp")
|
|
||||||
api.add_resource(ToolWorkflowListApi, "/workspaces/current/tools/workflow")
|
|
||||||
api.add_resource(ToolLabelsApi, "/workspaces/current/tool-labels")
|
|
||||||
|
|||||||
@ -1,7 +1,6 @@
|
|||||||
import logging
|
import logging
|
||||||
|
|
||||||
from flask import request
|
from flask import request
|
||||||
from flask_login import current_user
|
|
||||||
from flask_restx import Resource, fields, inputs, marshal, marshal_with, reqparse
|
from flask_restx import Resource, fields, inputs, marshal, marshal_with, reqparse
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from werkzeug.exceptions import Unauthorized
|
from werkzeug.exceptions import Unauthorized
|
||||||
@ -14,7 +13,7 @@ from controllers.common.errors import (
|
|||||||
TooManyFilesError,
|
TooManyFilesError,
|
||||||
UnsupportedFileTypeError,
|
UnsupportedFileTypeError,
|
||||||
)
|
)
|
||||||
from controllers.console import api
|
from controllers.console import console_ns
|
||||||
from controllers.console.admin import admin_required
|
from controllers.console.admin import admin_required
|
||||||
from controllers.console.error import AccountNotLinkTenantError
|
from controllers.console.error import AccountNotLinkTenantError
|
||||||
from controllers.console.wraps import (
|
from controllers.console.wraps import (
|
||||||
@ -24,7 +23,7 @@ from controllers.console.wraps import (
|
|||||||
)
|
)
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from libs.helper import TimestampField
|
from libs.helper import TimestampField
|
||||||
from libs.login import login_required
|
from libs.login import current_user, login_required
|
||||||
from models.account import Account, Tenant, TenantStatus
|
from models.account import Account, Tenant, TenantStatus
|
||||||
from services.account_service import TenantService
|
from services.account_service import TenantService
|
||||||
from services.feature_service import FeatureService
|
from services.feature_service import FeatureService
|
||||||
@ -65,6 +64,7 @@ tenants_fields = {
|
|||||||
workspace_fields = {"id": fields.String, "name": fields.String, "status": fields.String, "created_at": TimestampField}
|
workspace_fields = {"id": fields.String, "name": fields.String, "status": fields.String, "created_at": TimestampField}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces")
|
||||||
class TenantListApi(Resource):
|
class TenantListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -93,6 +93,7 @@ class TenantListApi(Resource):
|
|||||||
return {"workspaces": marshal(tenant_dicts, tenants_fields)}, 200
|
return {"workspaces": marshal(tenant_dicts, tenants_fields)}, 200
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/all-workspaces")
|
||||||
class WorkspaceListApi(Resource):
|
class WorkspaceListApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@admin_required
|
@admin_required
|
||||||
@ -118,6 +119,8 @@ class WorkspaceListApi(Resource):
|
|||||||
}, 200
|
}, 200
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/current", endpoint="workspaces_current")
|
||||||
|
@console_ns.route("/info", endpoint="info") # Deprecated
|
||||||
class TenantApi(Resource):
|
class TenantApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -143,11 +146,10 @@ class TenantApi(Resource):
|
|||||||
else:
|
else:
|
||||||
raise Unauthorized("workspace is archived")
|
raise Unauthorized("workspace is archived")
|
||||||
|
|
||||||
if not tenant:
|
|
||||||
raise ValueError("No tenant available")
|
|
||||||
return WorkspaceService.get_tenant_info(tenant), 200
|
return WorkspaceService.get_tenant_info(tenant), 200
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/switch")
|
||||||
class SwitchWorkspaceApi(Resource):
|
class SwitchWorkspaceApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -172,6 +174,7 @@ class SwitchWorkspaceApi(Resource):
|
|||||||
return {"result": "success", "new_tenant": marshal(WorkspaceService.get_tenant_info(new_tenant), tenant_fields)}
|
return {"result": "success", "new_tenant": marshal(WorkspaceService.get_tenant_info(new_tenant), tenant_fields)}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/custom-config")
|
||||||
class CustomConfigWorkspaceApi(Resource):
|
class CustomConfigWorkspaceApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -202,6 +205,7 @@ class CustomConfigWorkspaceApi(Resource):
|
|||||||
return {"result": "success", "tenant": marshal(WorkspaceService.get_tenant_info(tenant), tenant_fields)}
|
return {"result": "success", "tenant": marshal(WorkspaceService.get_tenant_info(tenant), tenant_fields)}
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/custom-config/webapp-logo/upload")
|
||||||
class WebappLogoWorkspaceApi(Resource):
|
class WebappLogoWorkspaceApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -242,6 +246,7 @@ class WebappLogoWorkspaceApi(Resource):
|
|||||||
return {"id": upload_file.id}, 201
|
return {"id": upload_file.id}, 201
|
||||||
|
|
||||||
|
|
||||||
|
@console_ns.route("/workspaces/info")
|
||||||
class WorkspaceInfoApi(Resource):
|
class WorkspaceInfoApi(Resource):
|
||||||
@setup_required
|
@setup_required
|
||||||
@login_required
|
@login_required
|
||||||
@ -261,13 +266,3 @@ class WorkspaceInfoApi(Resource):
|
|||||||
db.session.commit()
|
db.session.commit()
|
||||||
|
|
||||||
return {"result": "success", "tenant": marshal(WorkspaceService.get_tenant_info(tenant), tenant_fields)}
|
return {"result": "success", "tenant": marshal(WorkspaceService.get_tenant_info(tenant), tenant_fields)}
|
||||||
|
|
||||||
|
|
||||||
api.add_resource(TenantListApi, "/workspaces") # GET for getting all tenants
|
|
||||||
api.add_resource(WorkspaceListApi, "/all-workspaces") # GET for getting all tenants
|
|
||||||
api.add_resource(TenantApi, "/workspaces/current", endpoint="workspaces_current") # GET for getting current tenant info
|
|
||||||
api.add_resource(TenantApi, "/info", endpoint="info") # Deprecated
|
|
||||||
api.add_resource(SwitchWorkspaceApi, "/workspaces/switch") # POST for switching tenant
|
|
||||||
api.add_resource(CustomConfigWorkspaceApi, "/workspaces/custom-config")
|
|
||||||
api.add_resource(WebappLogoWorkspaceApi, "/workspaces/custom-config/webapp-logo/upload")
|
|
||||||
api.add_resource(WorkspaceInfoApi, "/workspaces/info") # POST for changing workspace info
|
|
||||||
|
|||||||
@ -7,13 +7,13 @@ from functools import wraps
|
|||||||
from typing import ParamSpec, TypeVar
|
from typing import ParamSpec, TypeVar
|
||||||
|
|
||||||
from flask import abort, request
|
from flask import abort, request
|
||||||
from flask_login import current_user
|
|
||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from controllers.console.workspace.error import AccountNotInitializedError
|
from controllers.console.workspace.error import AccountNotInitializedError
|
||||||
from extensions.ext_database import db
|
from extensions.ext_database import db
|
||||||
from extensions.ext_redis import redis_client
|
from extensions.ext_redis import redis_client
|
||||||
from models.account import AccountStatus
|
from libs.login import current_user
|
||||||
|
from models.account import Account, AccountStatus
|
||||||
from models.dataset import RateLimitLog
|
from models.dataset import RateLimitLog
|
||||||
from models.model import DifySetup
|
from models.model import DifySetup
|
||||||
from services.feature_service import FeatureService, LicenseStatus
|
from services.feature_service import FeatureService, LicenseStatus
|
||||||
@ -25,11 +25,16 @@ P = ParamSpec("P")
|
|||||||
R = TypeVar("R")
|
R = TypeVar("R")
|
||||||
|
|
||||||
|
|
||||||
|
def _current_account() -> Account:
|
||||||
|
assert isinstance(current_user, Account)
|
||||||
|
return current_user
|
||||||
|
|
||||||
|
|
||||||
def account_initialization_required(view: Callable[P, R]):
|
def account_initialization_required(view: Callable[P, R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
# check account initialization
|
# check account initialization
|
||||||
account = current_user
|
account = _current_account()
|
||||||
|
|
||||||
if account.status == AccountStatus.UNINITIALIZED:
|
if account.status == AccountStatus.UNINITIALIZED:
|
||||||
raise AccountNotInitializedError()
|
raise AccountNotInitializedError()
|
||||||
@ -75,7 +80,9 @@ def only_edition_self_hosted(view: Callable[P, R]):
|
|||||||
def cloud_edition_billing_enabled(view: Callable[P, R]):
|
def cloud_edition_billing_enabled(view: Callable[P, R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
features = FeatureService.get_features(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
features = FeatureService.get_features(account.current_tenant_id)
|
||||||
if not features.billing.enabled:
|
if not features.billing.enabled:
|
||||||
abort(403, "Billing feature is not enabled.")
|
abort(403, "Billing feature is not enabled.")
|
||||||
return view(*args, **kwargs)
|
return view(*args, **kwargs)
|
||||||
@ -87,7 +94,10 @@ def cloud_edition_billing_resource_check(resource: str):
|
|||||||
def interceptor(view: Callable[P, R]):
|
def interceptor(view: Callable[P, R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
features = FeatureService.get_features(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
tenant_id = account.current_tenant_id
|
||||||
|
features = FeatureService.get_features(tenant_id)
|
||||||
if features.billing.enabled:
|
if features.billing.enabled:
|
||||||
members = features.members
|
members = features.members
|
||||||
apps = features.apps
|
apps = features.apps
|
||||||
@ -128,7 +138,9 @@ def cloud_edition_billing_knowledge_limit_check(resource: str):
|
|||||||
def interceptor(view: Callable[P, R]):
|
def interceptor(view: Callable[P, R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
features = FeatureService.get_features(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
features = FeatureService.get_features(account.current_tenant_id)
|
||||||
if features.billing.enabled:
|
if features.billing.enabled:
|
||||||
if resource == "add_segment":
|
if resource == "add_segment":
|
||||||
if features.billing.subscription.plan == "sandbox":
|
if features.billing.subscription.plan == "sandbox":
|
||||||
@ -151,10 +163,13 @@ def cloud_edition_billing_rate_limit_check(resource: str):
|
|||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
if resource == "knowledge":
|
if resource == "knowledge":
|
||||||
knowledge_rate_limit = FeatureService.get_knowledge_rate_limit(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
tenant_id = account.current_tenant_id
|
||||||
|
knowledge_rate_limit = FeatureService.get_knowledge_rate_limit(tenant_id)
|
||||||
if knowledge_rate_limit.enabled:
|
if knowledge_rate_limit.enabled:
|
||||||
current_time = int(time.time() * 1000)
|
current_time = int(time.time() * 1000)
|
||||||
key = f"rate_limit_{current_user.current_tenant_id}"
|
key = f"rate_limit_{tenant_id}"
|
||||||
|
|
||||||
redis_client.zadd(key, {current_time: current_time})
|
redis_client.zadd(key, {current_time: current_time})
|
||||||
|
|
||||||
@ -165,7 +180,7 @@ def cloud_edition_billing_rate_limit_check(resource: str):
|
|||||||
if request_count > knowledge_rate_limit.limit:
|
if request_count > knowledge_rate_limit.limit:
|
||||||
# add ratelimit record
|
# add ratelimit record
|
||||||
rate_limit_log = RateLimitLog(
|
rate_limit_log = RateLimitLog(
|
||||||
tenant_id=current_user.current_tenant_id,
|
tenant_id=tenant_id,
|
||||||
subscription_plan=knowledge_rate_limit.subscription_plan,
|
subscription_plan=knowledge_rate_limit.subscription_plan,
|
||||||
operation="knowledge",
|
operation="knowledge",
|
||||||
)
|
)
|
||||||
@ -185,14 +200,17 @@ def cloud_utm_record(view: Callable[P, R]):
|
|||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
with contextlib.suppress(Exception):
|
with contextlib.suppress(Exception):
|
||||||
features = FeatureService.get_features(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
tenant_id = account.current_tenant_id
|
||||||
|
features = FeatureService.get_features(tenant_id)
|
||||||
|
|
||||||
if features.billing.enabled:
|
if features.billing.enabled:
|
||||||
utm_info = request.cookies.get("utm_info")
|
utm_info = request.cookies.get("utm_info")
|
||||||
|
|
||||||
if utm_info:
|
if utm_info:
|
||||||
utm_info_dict: dict = json.loads(utm_info)
|
utm_info_dict: dict = json.loads(utm_info)
|
||||||
OperationService.record_utm(current_user.current_tenant_id, utm_info_dict)
|
OperationService.record_utm(tenant_id, utm_info_dict)
|
||||||
|
|
||||||
return view(*args, **kwargs)
|
return view(*args, **kwargs)
|
||||||
|
|
||||||
@ -271,7 +289,9 @@ def enable_change_email(view: Callable[P, R]):
|
|||||||
def is_allow_transfer_owner(view: Callable[P, R]):
|
def is_allow_transfer_owner(view: Callable[P, R]):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args: P.args, **kwargs: P.kwargs):
|
def decorated(*args: P.args, **kwargs: P.kwargs):
|
||||||
features = FeatureService.get_features(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
features = FeatureService.get_features(account.current_tenant_id)
|
||||||
if features.is_allow_transfer_workspace:
|
if features.is_allow_transfer_workspace:
|
||||||
return view(*args, **kwargs)
|
return view(*args, **kwargs)
|
||||||
|
|
||||||
@ -284,7 +304,9 @@ def is_allow_transfer_owner(view: Callable[P, R]):
|
|||||||
def knowledge_pipeline_publish_enabled(view):
|
def knowledge_pipeline_publish_enabled(view):
|
||||||
@wraps(view)
|
@wraps(view)
|
||||||
def decorated(*args, **kwargs):
|
def decorated(*args, **kwargs):
|
||||||
features = FeatureService.get_features(current_user.current_tenant_id)
|
account = _current_account()
|
||||||
|
assert account.current_tenant_id is not None
|
||||||
|
features = FeatureService.get_features(account.current_tenant_id)
|
||||||
if features.knowledge_pipeline.publish_enabled:
|
if features.knowledge_pipeline.publish_enabled:
|
||||||
return view(*args, **kwargs)
|
return view(*args, **kwargs)
|
||||||
abort(403)
|
abort(403)
|
||||||
|
|||||||
@ -25,8 +25,8 @@ def get_user(tenant_id: str, user_id: str | None) -> EndUser:
|
|||||||
As a result, it could only be considered as an end user id.
|
As a result, it could only be considered as an end user id.
|
||||||
"""
|
"""
|
||||||
if not user_id:
|
if not user_id:
|
||||||
user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID.value
|
user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID
|
||||||
is_anonymous = user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID.value
|
is_anonymous = user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID
|
||||||
try:
|
try:
|
||||||
with Session(db.engine) as session:
|
with Session(db.engine) as session:
|
||||||
user_model = None
|
user_model = None
|
||||||
@ -85,7 +85,7 @@ def get_user_tenant(view: Callable[P, R] | None = None):
|
|||||||
raise ValueError("tenant_id is required")
|
raise ValueError("tenant_id is required")
|
||||||
|
|
||||||
if not user_id:
|
if not user_id:
|
||||||
user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID.value
|
user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID
|
||||||
|
|
||||||
try:
|
try:
|
||||||
tenant_model = (
|
tenant_model = (
|
||||||
@ -128,7 +128,7 @@ def plugin_data(view: Callable[P, R] | None = None, *, payload_type: type[BaseMo
|
|||||||
raise ValueError("invalid json")
|
raise ValueError("invalid json")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
payload = payload_type(**data)
|
payload = payload_type.model_validate(data)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise ValueError(f"invalid payload: {str(e)}")
|
raise ValueError(f"invalid payload: {str(e)}")
|
||||||
|
|
||||||
|
|||||||
@ -280,7 +280,7 @@ class DatasetListApi(DatasetApiResource):
|
|||||||
external_knowledge_id=args["external_knowledge_id"],
|
external_knowledge_id=args["external_knowledge_id"],
|
||||||
embedding_model_provider=args["embedding_model_provider"],
|
embedding_model_provider=args["embedding_model_provider"],
|
||||||
embedding_model_name=args["embedding_model"],
|
embedding_model_name=args["embedding_model"],
|
||||||
retrieval_model=RetrievalModel(**args["retrieval_model"])
|
retrieval_model=RetrievalModel.model_validate(args["retrieval_model"])
|
||||||
if args["retrieval_model"] is not None
|
if args["retrieval_model"] is not None
|
||||||
else None,
|
else None,
|
||||||
)
|
)
|
||||||
|
|||||||
@ -136,7 +136,7 @@ class DocumentAddByTextApi(DatasetApiResource):
|
|||||||
"info_list": {"data_source_type": "upload_file", "file_info_list": {"file_ids": [upload_file.id]}},
|
"info_list": {"data_source_type": "upload_file", "file_info_list": {"file_ids": [upload_file.id]}},
|
||||||
}
|
}
|
||||||
args["data_source"] = data_source
|
args["data_source"] = data_source
|
||||||
knowledge_config = KnowledgeConfig(**args)
|
knowledge_config = KnowledgeConfig.model_validate(args)
|
||||||
# validate args
|
# validate args
|
||||||
DocumentService.document_create_args_validate(knowledge_config)
|
DocumentService.document_create_args_validate(knowledge_config)
|
||||||
|
|
||||||
@ -221,7 +221,7 @@ class DocumentUpdateByTextApi(DatasetApiResource):
|
|||||||
args["data_source"] = data_source
|
args["data_source"] = data_source
|
||||||
# validate args
|
# validate args
|
||||||
args["original_document_id"] = str(document_id)
|
args["original_document_id"] = str(document_id)
|
||||||
knowledge_config = KnowledgeConfig(**args)
|
knowledge_config = KnowledgeConfig.model_validate(args)
|
||||||
DocumentService.document_create_args_validate(knowledge_config)
|
DocumentService.document_create_args_validate(knowledge_config)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@ -328,7 +328,7 @@ class DocumentAddByFileApi(DatasetApiResource):
|
|||||||
}
|
}
|
||||||
args["data_source"] = data_source
|
args["data_source"] = data_source
|
||||||
# validate args
|
# validate args
|
||||||
knowledge_config = KnowledgeConfig(**args)
|
knowledge_config = KnowledgeConfig.model_validate(args)
|
||||||
DocumentService.document_create_args_validate(knowledge_config)
|
DocumentService.document_create_args_validate(knowledge_config)
|
||||||
|
|
||||||
dataset_process_rule = dataset.latest_process_rule if "process_rule" not in args else None
|
dataset_process_rule = dataset.latest_process_rule if "process_rule" not in args else None
|
||||||
@ -426,7 +426,7 @@ class DocumentUpdateByFileApi(DatasetApiResource):
|
|||||||
# validate args
|
# validate args
|
||||||
args["original_document_id"] = str(document_id)
|
args["original_document_id"] = str(document_id)
|
||||||
|
|
||||||
knowledge_config = KnowledgeConfig(**args)
|
knowledge_config = KnowledgeConfig.model_validate(args)
|
||||||
DocumentService.document_create_args_validate(knowledge_config)
|
DocumentService.document_create_args_validate(knowledge_config)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|||||||
@ -51,7 +51,7 @@ class DatasetMetadataCreateServiceApi(DatasetApiResource):
|
|||||||
def post(self, tenant_id, dataset_id):
|
def post(self, tenant_id, dataset_id):
|
||||||
"""Create metadata for a dataset."""
|
"""Create metadata for a dataset."""
|
||||||
args = metadata_create_parser.parse_args()
|
args = metadata_create_parser.parse_args()
|
||||||
metadata_args = MetadataArgs(**args)
|
metadata_args = MetadataArgs.model_validate(args)
|
||||||
|
|
||||||
dataset_id_str = str(dataset_id)
|
dataset_id_str = str(dataset_id)
|
||||||
dataset = DatasetService.get_dataset(dataset_id_str)
|
dataset = DatasetService.get_dataset(dataset_id_str)
|
||||||
@ -200,7 +200,7 @@ class DocumentMetadataEditServiceApi(DatasetApiResource):
|
|||||||
DatasetService.check_dataset_permission(dataset, current_user)
|
DatasetService.check_dataset_permission(dataset, current_user)
|
||||||
|
|
||||||
args = document_metadata_parser.parse_args()
|
args = document_metadata_parser.parse_args()
|
||||||
metadata_args = MetadataOperationData(**args)
|
metadata_args = MetadataOperationData.model_validate(args)
|
||||||
|
|
||||||
MetadataService.update_documents_metadata(dataset, metadata_args)
|
MetadataService.update_documents_metadata(dataset, metadata_args)
|
||||||
|
|
||||||
|
|||||||
@ -98,7 +98,7 @@ class DatasourceNodeRunApi(DatasetApiResource):
|
|||||||
parser.add_argument("is_published", type=bool, required=True, location="json")
|
parser.add_argument("is_published", type=bool, required=True, location="json")
|
||||||
args: ParseResult = parser.parse_args()
|
args: ParseResult = parser.parse_args()
|
||||||
|
|
||||||
datasource_node_run_api_entity: DatasourceNodeRunApiEntity = DatasourceNodeRunApiEntity(**args)
|
datasource_node_run_api_entity = DatasourceNodeRunApiEntity.model_validate(args)
|
||||||
assert isinstance(current_user, Account)
|
assert isinstance(current_user, Account)
|
||||||
rag_pipeline_service: RagPipelineService = RagPipelineService()
|
rag_pipeline_service: RagPipelineService = RagPipelineService()
|
||||||
pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id)
|
pipeline: Pipeline = rag_pipeline_service.get_pipeline(tenant_id=tenant_id, dataset_id=dataset_id)
|
||||||
|
|||||||
@ -252,7 +252,7 @@ class DatasetSegmentApi(DatasetApiResource):
|
|||||||
args = segment_update_parser.parse_args()
|
args = segment_update_parser.parse_args()
|
||||||
|
|
||||||
updated_segment = SegmentService.update_segment(
|
updated_segment = SegmentService.update_segment(
|
||||||
SegmentUpdateArgs(**args["segment"]), segment, document, dataset
|
SegmentUpdateArgs.model_validate(args["segment"]), segment, document, dataset
|
||||||
)
|
)
|
||||||
return {"data": marshal(updated_segment, segment_fields), "doc_form": document.doc_form}, 200
|
return {"data": marshal(updated_segment, segment_fields), "doc_form": document.doc_form}, 200
|
||||||
|
|
||||||
|
|||||||
@ -313,7 +313,7 @@ def create_or_update_end_user_for_user_id(app_model: App, user_id: str | None =
|
|||||||
Create or update session terminal based on user ID.
|
Create or update session terminal based on user ID.
|
||||||
"""
|
"""
|
||||||
if not user_id:
|
if not user_id:
|
||||||
user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID.value
|
user_id = DefaultEndUserSessionID.DEFAULT_SESSION_ID
|
||||||
|
|
||||||
with Session(db.engine, expire_on_commit=False) as session:
|
with Session(db.engine, expire_on_commit=False) as session:
|
||||||
end_user = (
|
end_user = (
|
||||||
@ -332,7 +332,7 @@ def create_or_update_end_user_for_user_id(app_model: App, user_id: str | None =
|
|||||||
tenant_id=app_model.tenant_id,
|
tenant_id=app_model.tenant_id,
|
||||||
app_id=app_model.id,
|
app_id=app_model.id,
|
||||||
type="service_api",
|
type="service_api",
|
||||||
is_anonymous=user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID.value,
|
is_anonymous=user_id == DefaultEndUserSessionID.DEFAULT_SESSION_ID,
|
||||||
session_id=user_id,
|
session_id=user_id,
|
||||||
)
|
)
|
||||||
session.add(end_user)
|
session.add(end_user)
|
||||||
|
|||||||
@ -126,6 +126,8 @@ def exchange_token_for_existing_web_user(app_code: str, enterprise_user_decoded:
|
|||||||
end_user_id = enterprise_user_decoded.get("end_user_id")
|
end_user_id = enterprise_user_decoded.get("end_user_id")
|
||||||
session_id = enterprise_user_decoded.get("session_id")
|
session_id = enterprise_user_decoded.get("session_id")
|
||||||
user_auth_type = enterprise_user_decoded.get("auth_type")
|
user_auth_type = enterprise_user_decoded.get("auth_type")
|
||||||
|
exchanged_token_expires_unix = enterprise_user_decoded.get("exp")
|
||||||
|
|
||||||
if not user_auth_type:
|
if not user_auth_type:
|
||||||
raise Unauthorized("Missing auth_type in the token.")
|
raise Unauthorized("Missing auth_type in the token.")
|
||||||
|
|
||||||
@ -169,8 +171,11 @@ def exchange_token_for_existing_web_user(app_code: str, enterprise_user_decoded:
|
|||||||
)
|
)
|
||||||
db.session.add(end_user)
|
db.session.add(end_user)
|
||||||
db.session.commit()
|
db.session.commit()
|
||||||
exp_dt = datetime.now(UTC) + timedelta(minutes=dify_config.ACCESS_TOKEN_EXPIRE_MINUTES)
|
|
||||||
exp = int(exp_dt.timestamp())
|
exp = int((datetime.now(UTC) + timedelta(minutes=dify_config.ACCESS_TOKEN_EXPIRE_MINUTES)).timestamp())
|
||||||
|
if exchanged_token_expires_unix:
|
||||||
|
exp = int(exchanged_token_expires_unix)
|
||||||
|
|
||||||
payload = {
|
payload = {
|
||||||
"iss": site.id,
|
"iss": site.id,
|
||||||
"sub": "Web API Passport",
|
"sub": "Web API Passport",
|
||||||
|
|||||||
@ -40,7 +40,7 @@ class AgentConfigManager:
|
|||||||
"credential_id": tool.get("credential_id", None),
|
"credential_id": tool.get("credential_id", None),
|
||||||
}
|
}
|
||||||
|
|
||||||
agent_tools.append(AgentToolEntity(**agent_tool_properties))
|
agent_tools.append(AgentToolEntity.model_validate(agent_tool_properties))
|
||||||
|
|
||||||
if "strategy" in config["agent_mode"] and config["agent_mode"]["strategy"] not in {
|
if "strategy" in config["agent_mode"] and config["agent_mode"]["strategy"] not in {
|
||||||
"react_router",
|
"react_router",
|
||||||
|
|||||||
@ -197,12 +197,12 @@ class DatasetConfigManager:
|
|||||||
|
|
||||||
# strategy
|
# strategy
|
||||||
if "strategy" not in config["agent_mode"] or not config["agent_mode"].get("strategy"):
|
if "strategy" not in config["agent_mode"] or not config["agent_mode"].get("strategy"):
|
||||||
config["agent_mode"]["strategy"] = PlanningStrategy.ROUTER.value
|
config["agent_mode"]["strategy"] = PlanningStrategy.ROUTER
|
||||||
|
|
||||||
has_datasets = False
|
has_datasets = False
|
||||||
if config.get("agent_mode", {}).get("strategy") in {
|
if config.get("agent_mode", {}).get("strategy") in {
|
||||||
PlanningStrategy.ROUTER.value,
|
PlanningStrategy.ROUTER,
|
||||||
PlanningStrategy.REACT_ROUTER.value,
|
PlanningStrategy.REACT_ROUTER,
|
||||||
}:
|
}:
|
||||||
for tool in config.get("agent_mode", {}).get("tools", []):
|
for tool in config.get("agent_mode", {}).get("tools", []):
|
||||||
key = list(tool.keys())[0]
|
key = list(tool.keys())[0]
|
||||||
|
|||||||
@ -68,9 +68,13 @@ class ModelConfigConverter:
|
|||||||
# get model mode
|
# get model mode
|
||||||
model_mode = model_config.mode
|
model_mode = model_config.mode
|
||||||
if not model_mode:
|
if not model_mode:
|
||||||
model_mode = LLMMode.CHAT.value
|
model_mode = LLMMode.CHAT
|
||||||
if model_schema and model_schema.model_properties.get(ModelPropertyKey.MODE):
|
if model_schema and model_schema.model_properties.get(ModelPropertyKey.MODE):
|
||||||
model_mode = LLMMode(model_schema.model_properties[ModelPropertyKey.MODE]).value
|
try:
|
||||||
|
model_mode = LLMMode(model_schema.model_properties[ModelPropertyKey.MODE])
|
||||||
|
except ValueError:
|
||||||
|
# Fall back to CHAT mode if the stored value is invalid
|
||||||
|
model_mode = LLMMode.CHAT
|
||||||
|
|
||||||
if not model_schema:
|
if not model_schema:
|
||||||
raise ValueError(f"Model {model_name} not exist.")
|
raise ValueError(f"Model {model_name} not exist.")
|
||||||
|
|||||||
@ -100,7 +100,7 @@ class PromptTemplateConfigManager:
|
|||||||
if config["model"]["mode"] not in model_mode_vals:
|
if config["model"]["mode"] not in model_mode_vals:
|
||||||
raise ValueError(f"model.mode must be in {model_mode_vals} when prompt_type is advanced")
|
raise ValueError(f"model.mode must be in {model_mode_vals} when prompt_type is advanced")
|
||||||
|
|
||||||
if app_mode == AppMode.CHAT and config["model"]["mode"] == ModelMode.COMPLETION.value:
|
if app_mode == AppMode.CHAT and config["model"]["mode"] == ModelMode.COMPLETION:
|
||||||
user_prefix = config["completion_prompt_config"]["conversation_histories_role"]["user_prefix"]
|
user_prefix = config["completion_prompt_config"]["conversation_histories_role"]["user_prefix"]
|
||||||
assistant_prefix = config["completion_prompt_config"]["conversation_histories_role"]["assistant_prefix"]
|
assistant_prefix = config["completion_prompt_config"]["conversation_histories_role"]["assistant_prefix"]
|
||||||
|
|
||||||
@ -110,7 +110,7 @@ class PromptTemplateConfigManager:
|
|||||||
if not assistant_prefix:
|
if not assistant_prefix:
|
||||||
config["completion_prompt_config"]["conversation_histories_role"]["assistant_prefix"] = "Assistant"
|
config["completion_prompt_config"]["conversation_histories_role"]["assistant_prefix"] = "Assistant"
|
||||||
|
|
||||||
if config["model"]["mode"] == ModelMode.CHAT.value:
|
if config["model"]["mode"] == ModelMode.CHAT:
|
||||||
prompt_list = config["chat_prompt_config"]["prompt"]
|
prompt_list = config["chat_prompt_config"]["prompt"]
|
||||||
|
|
||||||
if len(prompt_list) > 10:
|
if len(prompt_list) > 10:
|
||||||
|
|||||||
@ -186,7 +186,7 @@ class AgentChatAppConfigManager(BaseAppConfigManager):
|
|||||||
raise ValueError("enabled in agent_mode must be of boolean type")
|
raise ValueError("enabled in agent_mode must be of boolean type")
|
||||||
|
|
||||||
if not agent_mode.get("strategy"):
|
if not agent_mode.get("strategy"):
|
||||||
agent_mode["strategy"] = PlanningStrategy.ROUTER.value
|
agent_mode["strategy"] = PlanningStrategy.ROUTER
|
||||||
|
|
||||||
if agent_mode["strategy"] not in [member.value for member in list(PlanningStrategy.__members__.values())]:
|
if agent_mode["strategy"] not in [member.value for member in list(PlanningStrategy.__members__.values())]:
|
||||||
raise ValueError("strategy in agent_mode must be in the specified strategy list")
|
raise ValueError("strategy in agent_mode must be in the specified strategy list")
|
||||||
|
|||||||
@ -198,9 +198,9 @@ class AgentChatAppRunner(AppRunner):
|
|||||||
# start agent runner
|
# start agent runner
|
||||||
if agent_entity.strategy == AgentEntity.Strategy.CHAIN_OF_THOUGHT:
|
if agent_entity.strategy == AgentEntity.Strategy.CHAIN_OF_THOUGHT:
|
||||||
# check LLM mode
|
# check LLM mode
|
||||||
if model_schema.model_properties.get(ModelPropertyKey.MODE) == LLMMode.CHAT.value:
|
if model_schema.model_properties.get(ModelPropertyKey.MODE) == LLMMode.CHAT:
|
||||||
runner_cls = CotChatAgentRunner
|
runner_cls = CotChatAgentRunner
|
||||||
elif model_schema.model_properties.get(ModelPropertyKey.MODE) == LLMMode.COMPLETION.value:
|
elif model_schema.model_properties.get(ModelPropertyKey.MODE) == LLMMode.COMPLETION:
|
||||||
runner_cls = CotCompletionAgentRunner
|
runner_cls = CotCompletionAgentRunner
|
||||||
else:
|
else:
|
||||||
raise ValueError(f"Invalid LLM mode: {model_schema.model_properties.get(ModelPropertyKey.MODE)}")
|
raise ValueError(f"Invalid LLM mode: {model_schema.model_properties.get(ModelPropertyKey.MODE)}")
|
||||||
|
|||||||
@ -61,9 +61,6 @@ class AppRunner:
|
|||||||
if model_context_tokens is None:
|
if model_context_tokens is None:
|
||||||
return -1
|
return -1
|
||||||
|
|
||||||
if max_tokens is None:
|
|
||||||
max_tokens = 0
|
|
||||||
|
|
||||||
prompt_tokens = model_instance.get_llm_num_tokens(prompt_messages)
|
prompt_tokens = model_instance.get_llm_num_tokens(prompt_messages)
|
||||||
|
|
||||||
if prompt_tokens + max_tokens > model_context_tokens:
|
if prompt_tokens + max_tokens > model_context_tokens:
|
||||||
|
|||||||
@ -116,7 +116,7 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
|||||||
rag_pipeline_variables = []
|
rag_pipeline_variables = []
|
||||||
if workflow.rag_pipeline_variables:
|
if workflow.rag_pipeline_variables:
|
||||||
for v in workflow.rag_pipeline_variables:
|
for v in workflow.rag_pipeline_variables:
|
||||||
rag_pipeline_variable = RAGPipelineVariable(**v)
|
rag_pipeline_variable = RAGPipelineVariable.model_validate(v)
|
||||||
if (
|
if (
|
||||||
rag_pipeline_variable.belong_to_node_id
|
rag_pipeline_variable.belong_to_node_id
|
||||||
in (self.application_generate_entity.start_node_id, "shared")
|
in (self.application_generate_entity.start_node_id, "shared")
|
||||||
@ -229,8 +229,8 @@ class PipelineRunner(WorkflowBasedAppRunner):
|
|||||||
workflow_id=workflow.id,
|
workflow_id=workflow.id,
|
||||||
graph_config=graph_config,
|
graph_config=graph_config,
|
||||||
user_id=self.application_generate_entity.user_id,
|
user_id=self.application_generate_entity.user_id,
|
||||||
user_from=UserFrom.ACCOUNT.value,
|
user_from=UserFrom.ACCOUNT,
|
||||||
invoke_from=InvokeFrom.SERVICE_API.value,
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
call_depth=0,
|
call_depth=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -100,8 +100,8 @@ class WorkflowBasedAppRunner:
|
|||||||
workflow_id=workflow_id,
|
workflow_id=workflow_id,
|
||||||
graph_config=graph_config,
|
graph_config=graph_config,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
user_from=UserFrom.ACCOUNT.value,
|
user_from=UserFrom.ACCOUNT,
|
||||||
invoke_from=InvokeFrom.SERVICE_API.value,
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
call_depth=0,
|
call_depth=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -244,8 +244,8 @@ class WorkflowBasedAppRunner:
|
|||||||
workflow_id=workflow.id,
|
workflow_id=workflow.id,
|
||||||
graph_config=graph_config,
|
graph_config=graph_config,
|
||||||
user_id="",
|
user_id="",
|
||||||
user_from=UserFrom.ACCOUNT.value,
|
user_from=UserFrom.ACCOUNT,
|
||||||
invoke_from=InvokeFrom.SERVICE_API.value,
|
invoke_from=InvokeFrom.SERVICE_API,
|
||||||
call_depth=0,
|
call_depth=0,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -49,7 +49,7 @@ class DatasourceProviderApiEntity(BaseModel):
|
|||||||
for datasource in datasources:
|
for datasource in datasources:
|
||||||
if datasource.get("parameters"):
|
if datasource.get("parameters"):
|
||||||
for parameter in datasource.get("parameters"):
|
for parameter in datasource.get("parameters"):
|
||||||
if parameter.get("type") == DatasourceParameter.DatasourceParameterType.SYSTEM_FILES.value:
|
if parameter.get("type") == DatasourceParameter.DatasourceParameterType.SYSTEM_FILES:
|
||||||
parameter["type"] = "files"
|
parameter["type"] = "files"
|
||||||
# -------------
|
# -------------
|
||||||
|
|
||||||
|
|||||||
@ -1,4 +1,4 @@
|
|||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field, model_validator
|
||||||
|
|
||||||
|
|
||||||
class I18nObject(BaseModel):
|
class I18nObject(BaseModel):
|
||||||
@ -11,11 +11,12 @@ class I18nObject(BaseModel):
|
|||||||
pt_BR: str | None = Field(default=None)
|
pt_BR: str | None = Field(default=None)
|
||||||
ja_JP: str | None = Field(default=None)
|
ja_JP: str | None = Field(default=None)
|
||||||
|
|
||||||
def __init__(self, **data):
|
@model_validator(mode="after")
|
||||||
super().__init__(**data)
|
def _(self):
|
||||||
self.zh_Hans = self.zh_Hans or self.en_US
|
self.zh_Hans = self.zh_Hans or self.en_US
|
||||||
self.pt_BR = self.pt_BR or self.en_US
|
self.pt_BR = self.pt_BR or self.en_US
|
||||||
self.ja_JP = self.ja_JP or self.en_US
|
self.ja_JP = self.ja_JP or self.en_US
|
||||||
|
return self
|
||||||
|
|
||||||
def to_dict(self) -> dict:
|
def to_dict(self) -> dict:
|
||||||
return {"zh_Hans": self.zh_Hans, "en_US": self.en_US, "pt_BR": self.pt_BR, "ja_JP": self.ja_JP}
|
return {"zh_Hans": self.zh_Hans, "en_US": self.en_US, "pt_BR": self.pt_BR, "ja_JP": self.ja_JP}
|
||||||
|
|||||||
@ -1,5 +1,5 @@
|
|||||||
import enum
|
import enum
|
||||||
from enum import Enum
|
from enum import StrEnum
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from pydantic import BaseModel, Field, ValidationInfo, field_validator
|
from pydantic import BaseModel, Field, ValidationInfo, field_validator
|
||||||
@ -54,16 +54,16 @@ class DatasourceParameter(PluginParameter):
|
|||||||
removes TOOLS_SELECTOR from PluginParameterType
|
removes TOOLS_SELECTOR from PluginParameterType
|
||||||
"""
|
"""
|
||||||
|
|
||||||
STRING = PluginParameterType.STRING.value
|
STRING = PluginParameterType.STRING
|
||||||
NUMBER = PluginParameterType.NUMBER.value
|
NUMBER = PluginParameterType.NUMBER
|
||||||
BOOLEAN = PluginParameterType.BOOLEAN.value
|
BOOLEAN = PluginParameterType.BOOLEAN
|
||||||
SELECT = PluginParameterType.SELECT.value
|
SELECT = PluginParameterType.SELECT
|
||||||
SECRET_INPUT = PluginParameterType.SECRET_INPUT.value
|
SECRET_INPUT = PluginParameterType.SECRET_INPUT
|
||||||
FILE = PluginParameterType.FILE.value
|
FILE = PluginParameterType.FILE
|
||||||
FILES = PluginParameterType.FILES.value
|
FILES = PluginParameterType.FILES
|
||||||
|
|
||||||
# deprecated, should not use.
|
# deprecated, should not use.
|
||||||
SYSTEM_FILES = PluginParameterType.SYSTEM_FILES.value
|
SYSTEM_FILES = PluginParameterType.SYSTEM_FILES
|
||||||
|
|
||||||
def as_normal_type(self):
|
def as_normal_type(self):
|
||||||
return as_normal_type(self)
|
return as_normal_type(self)
|
||||||
@ -218,7 +218,7 @@ class DatasourceLabel(BaseModel):
|
|||||||
icon: str = Field(..., description="The icon of the tool")
|
icon: str = Field(..., description="The icon of the tool")
|
||||||
|
|
||||||
|
|
||||||
class DatasourceInvokeFrom(Enum):
|
class DatasourceInvokeFrom(StrEnum):
|
||||||
"""
|
"""
|
||||||
Enum class for datasource invoke
|
Enum class for datasource invoke
|
||||||
"""
|
"""
|
||||||
|
|||||||
@ -5,7 +5,7 @@ from collections import defaultdict
|
|||||||
from collections.abc import Iterator, Sequence
|
from collections.abc import Iterator, Sequence
|
||||||
from json import JSONDecodeError
|
from json import JSONDecodeError
|
||||||
|
|
||||||
from pydantic import BaseModel, ConfigDict, Field
|
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||||
from sqlalchemy import func, select
|
from sqlalchemy import func, select
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
@ -73,9 +73,8 @@ class ProviderConfiguration(BaseModel):
|
|||||||
# pydantic configs
|
# pydantic configs
|
||||||
model_config = ConfigDict(protected_namespaces=())
|
model_config = ConfigDict(protected_namespaces=())
|
||||||
|
|
||||||
def __init__(self, **data):
|
@model_validator(mode="after")
|
||||||
super().__init__(**data)
|
def _(self):
|
||||||
|
|
||||||
if self.provider.provider not in original_provider_configurate_methods:
|
if self.provider.provider not in original_provider_configurate_methods:
|
||||||
original_provider_configurate_methods[self.provider.provider] = []
|
original_provider_configurate_methods[self.provider.provider] = []
|
||||||
for configurate_method in self.provider.configurate_methods:
|
for configurate_method in self.provider.configurate_methods:
|
||||||
@ -90,6 +89,7 @@ class ProviderConfiguration(BaseModel):
|
|||||||
and ConfigurateMethod.PREDEFINED_MODEL not in self.provider.configurate_methods
|
and ConfigurateMethod.PREDEFINED_MODEL not in self.provider.configurate_methods
|
||||||
):
|
):
|
||||||
self.provider.configurate_methods.append(ConfigurateMethod.PREDEFINED_MODEL)
|
self.provider.configurate_methods.append(ConfigurateMethod.PREDEFINED_MODEL)
|
||||||
|
return self
|
||||||
|
|
||||||
def get_current_credentials(self, model_type: ModelType, model: str) -> dict | None:
|
def get_current_credentials(self, model_type: ModelType, model: str) -> dict | None:
|
||||||
"""
|
"""
|
||||||
@ -207,7 +207,7 @@ class ProviderConfiguration(BaseModel):
|
|||||||
"""
|
"""
|
||||||
stmt = select(Provider).where(
|
stmt = select(Provider).where(
|
||||||
Provider.tenant_id == self.tenant_id,
|
Provider.tenant_id == self.tenant_id,
|
||||||
Provider.provider_type == ProviderType.CUSTOM.value,
|
Provider.provider_type == ProviderType.CUSTOM,
|
||||||
Provider.provider_name.in_(self._get_provider_names()),
|
Provider.provider_name.in_(self._get_provider_names()),
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -458,7 +458,7 @@ class ProviderConfiguration(BaseModel):
|
|||||||
provider_record = Provider(
|
provider_record = Provider(
|
||||||
tenant_id=self.tenant_id,
|
tenant_id=self.tenant_id,
|
||||||
provider_name=self.provider.provider,
|
provider_name=self.provider.provider,
|
||||||
provider_type=ProviderType.CUSTOM.value,
|
provider_type=ProviderType.CUSTOM,
|
||||||
is_valid=True,
|
is_valid=True,
|
||||||
credential_id=new_record.id,
|
credential_id=new_record.id,
|
||||||
)
|
)
|
||||||
@ -1414,7 +1414,7 @@ class ProviderConfiguration(BaseModel):
|
|||||||
"""
|
"""
|
||||||
secret_input_form_variables = []
|
secret_input_form_variables = []
|
||||||
for credential_form_schema in credential_form_schemas:
|
for credential_form_schema in credential_form_schemas:
|
||||||
if credential_form_schema.type.value == FormType.SECRET_INPUT.value:
|
if credential_form_schema.type == FormType.SECRET_INPUT:
|
||||||
secret_input_form_variables.append(credential_form_schema.variable)
|
secret_input_form_variables.append(credential_form_schema.variable)
|
||||||
|
|
||||||
return secret_input_form_variables
|
return secret_input_form_variables
|
||||||
|
|||||||
@ -1,13 +1,13 @@
|
|||||||
from typing import cast
|
from typing import cast
|
||||||
|
|
||||||
import requests
|
import httpx
|
||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
from models.api_based_extension import APIBasedExtensionPoint
|
from models.api_based_extension import APIBasedExtensionPoint
|
||||||
|
|
||||||
|
|
||||||
class APIBasedExtensionRequestor:
|
class APIBasedExtensionRequestor:
|
||||||
timeout: tuple[int, int] = (5, 60)
|
timeout: httpx.Timeout = httpx.Timeout(60.0, connect=5.0)
|
||||||
"""timeout for request connect and read"""
|
"""timeout for request connect and read"""
|
||||||
|
|
||||||
def __init__(self, api_endpoint: str, api_key: str):
|
def __init__(self, api_endpoint: str, api_key: str):
|
||||||
@ -27,25 +27,23 @@ class APIBasedExtensionRequestor:
|
|||||||
url = self.api_endpoint
|
url = self.api_endpoint
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# proxy support for security
|
mounts: dict[str, httpx.BaseTransport] | None = None
|
||||||
proxies = None
|
|
||||||
if dify_config.SSRF_PROXY_HTTP_URL and dify_config.SSRF_PROXY_HTTPS_URL:
|
if dify_config.SSRF_PROXY_HTTP_URL and dify_config.SSRF_PROXY_HTTPS_URL:
|
||||||
proxies = {
|
mounts = {
|
||||||
"http": dify_config.SSRF_PROXY_HTTP_URL,
|
"http://": httpx.HTTPTransport(proxy=dify_config.SSRF_PROXY_HTTP_URL),
|
||||||
"https": dify_config.SSRF_PROXY_HTTPS_URL,
|
"https://": httpx.HTTPTransport(proxy=dify_config.SSRF_PROXY_HTTPS_URL),
|
||||||
}
|
}
|
||||||
|
|
||||||
response = requests.request(
|
with httpx.Client(mounts=mounts, timeout=self.timeout) as client:
|
||||||
method="POST",
|
response = client.request(
|
||||||
url=url,
|
method="POST",
|
||||||
json={"point": point.value, "params": params},
|
url=url,
|
||||||
headers=headers,
|
json={"point": point.value, "params": params},
|
||||||
timeout=self.timeout,
|
headers=headers,
|
||||||
proxies=proxies,
|
)
|
||||||
)
|
except httpx.TimeoutException:
|
||||||
except requests.Timeout:
|
|
||||||
raise ValueError("request timeout")
|
raise ValueError("request timeout")
|
||||||
except requests.ConnectionError:
|
except httpx.RequestError:
|
||||||
raise ValueError("request connection error")
|
raise ValueError("request connection error")
|
||||||
|
|
||||||
if response.status_code != 200:
|
if response.status_code != 200:
|
||||||
|
|||||||
@ -131,7 +131,7 @@ class CodeExecutor:
|
|||||||
if (code := response_data.get("code")) != 0:
|
if (code := response_data.get("code")) != 0:
|
||||||
raise CodeExecutionError(f"Got error code: {code}. Got error msg: {response_data.get('message')}")
|
raise CodeExecutionError(f"Got error code: {code}. Got error msg: {response_data.get('message')}")
|
||||||
|
|
||||||
response_code = CodeExecutionResponse(**response_data)
|
response_code = CodeExecutionResponse.model_validate(response_data)
|
||||||
|
|
||||||
if response_code.data.error:
|
if response_code.data.error:
|
||||||
raise CodeExecutionError(response_code.data.error)
|
raise CodeExecutionError(response_code.data.error)
|
||||||
|
|||||||
@ -26,7 +26,7 @@ def batch_fetch_plugin_manifests(plugin_ids: list[str]) -> Sequence[MarketplaceP
|
|||||||
response = httpx.post(url, json={"plugin_ids": plugin_ids}, headers={"X-Dify-Version": dify_config.project.version})
|
response = httpx.post(url, json={"plugin_ids": plugin_ids}, headers={"X-Dify-Version": dify_config.project.version})
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
|
|
||||||
return [MarketplacePluginDeclaration(**plugin) for plugin in response.json()["data"]["plugins"]]
|
return [MarketplacePluginDeclaration.model_validate(plugin) for plugin in response.json()["data"]["plugins"]]
|
||||||
|
|
||||||
|
|
||||||
def batch_fetch_plugin_manifests_ignore_deserialization_error(
|
def batch_fetch_plugin_manifests_ignore_deserialization_error(
|
||||||
@ -41,7 +41,7 @@ def batch_fetch_plugin_manifests_ignore_deserialization_error(
|
|||||||
result: list[MarketplacePluginDeclaration] = []
|
result: list[MarketplacePluginDeclaration] = []
|
||||||
for plugin in response.json()["data"]["plugins"]:
|
for plugin in response.json()["data"]["plugins"]:
|
||||||
try:
|
try:
|
||||||
result.append(MarketplacePluginDeclaration(**plugin))
|
result.append(MarketplacePluginDeclaration.model_validate(plugin))
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|||||||
@ -20,7 +20,7 @@ from core.rag.cleaner.clean_processor import CleanProcessor
|
|||||||
from core.rag.datasource.keyword.keyword_factory import Keyword
|
from core.rag.datasource.keyword.keyword_factory import Keyword
|
||||||
from core.rag.docstore.dataset_docstore import DatasetDocumentStore
|
from core.rag.docstore.dataset_docstore import DatasetDocumentStore
|
||||||
from core.rag.extractor.entity.datasource_type import DatasourceType
|
from core.rag.extractor.entity.datasource_type import DatasourceType
|
||||||
from core.rag.extractor.entity.extract_setting import ExtractSetting
|
from core.rag.extractor.entity.extract_setting import ExtractSetting, NotionInfo, WebsiteInfo
|
||||||
from core.rag.index_processor.constant.index_type import IndexType
|
from core.rag.index_processor.constant.index_type import IndexType
|
||||||
from core.rag.index_processor.index_processor_base import BaseIndexProcessor
|
from core.rag.index_processor.index_processor_base import BaseIndexProcessor
|
||||||
from core.rag.index_processor.index_processor_factory import IndexProcessorFactory
|
from core.rag.index_processor.index_processor_factory import IndexProcessorFactory
|
||||||
@ -343,7 +343,7 @@ class IndexingRunner:
|
|||||||
|
|
||||||
if file_detail:
|
if file_detail:
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.FILE.value,
|
datasource_type=DatasourceType.FILE,
|
||||||
upload_file=file_detail,
|
upload_file=file_detail,
|
||||||
document_model=dataset_document.doc_form,
|
document_model=dataset_document.doc_form,
|
||||||
)
|
)
|
||||||
@ -356,15 +356,17 @@ class IndexingRunner:
|
|||||||
):
|
):
|
||||||
raise ValueError("no notion import info found")
|
raise ValueError("no notion import info found")
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.NOTION.value,
|
datasource_type=DatasourceType.NOTION,
|
||||||
notion_info={
|
notion_info=NotionInfo.model_validate(
|
||||||
"credential_id": data_source_info["credential_id"],
|
{
|
||||||
"notion_workspace_id": data_source_info["notion_workspace_id"],
|
"credential_id": data_source_info["credential_id"],
|
||||||
"notion_obj_id": data_source_info["notion_page_id"],
|
"notion_workspace_id": data_source_info["notion_workspace_id"],
|
||||||
"notion_page_type": data_source_info["type"],
|
"notion_obj_id": data_source_info["notion_page_id"],
|
||||||
"document": dataset_document,
|
"notion_page_type": data_source_info["type"],
|
||||||
"tenant_id": dataset_document.tenant_id,
|
"document": dataset_document,
|
||||||
},
|
"tenant_id": dataset_document.tenant_id,
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=dataset_document.doc_form,
|
document_model=dataset_document.doc_form,
|
||||||
)
|
)
|
||||||
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
|
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
|
||||||
@ -377,15 +379,17 @@ class IndexingRunner:
|
|||||||
):
|
):
|
||||||
raise ValueError("no website import info found")
|
raise ValueError("no website import info found")
|
||||||
extract_setting = ExtractSetting(
|
extract_setting = ExtractSetting(
|
||||||
datasource_type=DatasourceType.WEBSITE.value,
|
datasource_type=DatasourceType.WEBSITE,
|
||||||
website_info={
|
website_info=WebsiteInfo.model_validate(
|
||||||
"provider": data_source_info["provider"],
|
{
|
||||||
"job_id": data_source_info["job_id"],
|
"provider": data_source_info["provider"],
|
||||||
"tenant_id": dataset_document.tenant_id,
|
"job_id": data_source_info["job_id"],
|
||||||
"url": data_source_info["url"],
|
"tenant_id": dataset_document.tenant_id,
|
||||||
"mode": data_source_info["mode"],
|
"url": data_source_info["url"],
|
||||||
"only_main_content": data_source_info["only_main_content"],
|
"mode": data_source_info["mode"],
|
||||||
},
|
"only_main_content": data_source_info["only_main_content"],
|
||||||
|
}
|
||||||
|
),
|
||||||
document_model=dataset_document.doc_form,
|
document_model=dataset_document.doc_form,
|
||||||
)
|
)
|
||||||
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
|
text_docs = index_processor.extract(extract_setting, process_rule_mode=process_rule["mode"])
|
||||||
|
|||||||
@ -224,8 +224,8 @@ def _handle_native_json_schema(
|
|||||||
|
|
||||||
# Set appropriate response format if required by the model
|
# Set appropriate response format if required by the model
|
||||||
for rule in rules:
|
for rule in rules:
|
||||||
if rule.name == "response_format" and ResponseFormat.JSON_SCHEMA.value in rule.options:
|
if rule.name == "response_format" and ResponseFormat.JSON_SCHEMA in rule.options:
|
||||||
model_parameters["response_format"] = ResponseFormat.JSON_SCHEMA.value
|
model_parameters["response_format"] = ResponseFormat.JSON_SCHEMA
|
||||||
|
|
||||||
return model_parameters
|
return model_parameters
|
||||||
|
|
||||||
@ -239,10 +239,10 @@ def _set_response_format(model_parameters: dict, rules: list):
|
|||||||
"""
|
"""
|
||||||
for rule in rules:
|
for rule in rules:
|
||||||
if rule.name == "response_format":
|
if rule.name == "response_format":
|
||||||
if ResponseFormat.JSON.value in rule.options:
|
if ResponseFormat.JSON in rule.options:
|
||||||
model_parameters["response_format"] = ResponseFormat.JSON.value
|
model_parameters["response_format"] = ResponseFormat.JSON
|
||||||
elif ResponseFormat.JSON_OBJECT.value in rule.options:
|
elif ResponseFormat.JSON_OBJECT in rule.options:
|
||||||
model_parameters["response_format"] = ResponseFormat.JSON_OBJECT.value
|
model_parameters["response_format"] = ResponseFormat.JSON_OBJECT
|
||||||
|
|
||||||
|
|
||||||
def _handle_prompt_based_schema(
|
def _handle_prompt_based_schema(
|
||||||
|
|||||||
@ -294,7 +294,7 @@ class ClientSession(
|
|||||||
method="completion/complete",
|
method="completion/complete",
|
||||||
params=types.CompleteRequestParams(
|
params=types.CompleteRequestParams(
|
||||||
ref=ref,
|
ref=ref,
|
||||||
argument=types.CompletionArgument(**argument),
|
argument=types.CompletionArgument.model_validate(argument),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
),
|
),
|
||||||
|
|||||||
@ -1,4 +1,4 @@
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel, model_validator
|
||||||
|
|
||||||
|
|
||||||
class I18nObject(BaseModel):
|
class I18nObject(BaseModel):
|
||||||
@ -9,7 +9,8 @@ class I18nObject(BaseModel):
|
|||||||
zh_Hans: str | None = None
|
zh_Hans: str | None = None
|
||||||
en_US: str
|
en_US: str
|
||||||
|
|
||||||
def __init__(self, **data):
|
@model_validator(mode="after")
|
||||||
super().__init__(**data)
|
def _(self):
|
||||||
if not self.zh_Hans:
|
if not self.zh_Hans:
|
||||||
self.zh_Hans = self.en_US
|
self.zh_Hans = self.en_US
|
||||||
|
return self
|
||||||
|
|||||||
@ -1,13 +1,13 @@
|
|||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from enum import Enum, StrEnum, auto
|
from enum import StrEnum, auto
|
||||||
|
|
||||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||||
|
|
||||||
from core.model_runtime.entities.common_entities import I18nObject
|
from core.model_runtime.entities.common_entities import I18nObject
|
||||||
from core.model_runtime.entities.model_entities import AIModelEntity, ModelType
|
from core.model_runtime.entities.model_entities import AIModelEntity, ModelType
|
||||||
|
|
||||||
|
|
||||||
class ConfigurateMethod(Enum):
|
class ConfigurateMethod(StrEnum):
|
||||||
"""
|
"""
|
||||||
Enum class for configurate method of provider model.
|
Enum class for configurate method of provider model.
|
||||||
"""
|
"""
|
||||||
@ -46,10 +46,11 @@ class FormOption(BaseModel):
|
|||||||
value: str
|
value: str
|
||||||
show_on: list[FormShowOnObject] = []
|
show_on: list[FormShowOnObject] = []
|
||||||
|
|
||||||
def __init__(self, **data):
|
@model_validator(mode="after")
|
||||||
super().__init__(**data)
|
def _(self):
|
||||||
if not self.label:
|
if not self.label:
|
||||||
self.label = I18nObject(en_US=self.value)
|
self.label = I18nObject(en_US=self.value)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class CredentialFormSchema(BaseModel):
|
class CredentialFormSchema(BaseModel):
|
||||||
|
|||||||
@ -269,17 +269,17 @@ class ModelProviderFactory:
|
|||||||
}
|
}
|
||||||
|
|
||||||
if model_type == ModelType.LLM:
|
if model_type == ModelType.LLM:
|
||||||
return LargeLanguageModel(**init_params) # type: ignore
|
return LargeLanguageModel.model_validate(init_params)
|
||||||
elif model_type == ModelType.TEXT_EMBEDDING:
|
elif model_type == ModelType.TEXT_EMBEDDING:
|
||||||
return TextEmbeddingModel(**init_params) # type: ignore
|
return TextEmbeddingModel.model_validate(init_params)
|
||||||
elif model_type == ModelType.RERANK:
|
elif model_type == ModelType.RERANK:
|
||||||
return RerankModel(**init_params) # type: ignore
|
return RerankModel.model_validate(init_params)
|
||||||
elif model_type == ModelType.SPEECH2TEXT:
|
elif model_type == ModelType.SPEECH2TEXT:
|
||||||
return Speech2TextModel(**init_params) # type: ignore
|
return Speech2TextModel.model_validate(init_params)
|
||||||
elif model_type == ModelType.MODERATION:
|
elif model_type == ModelType.MODERATION:
|
||||||
return ModerationModel(**init_params) # type: ignore
|
return ModerationModel.model_validate(init_params)
|
||||||
elif model_type == ModelType.TTS:
|
elif model_type == ModelType.TTS:
|
||||||
return TTSModel(**init_params) # type: ignore
|
return TTSModel.model_validate(init_params)
|
||||||
|
|
||||||
def get_provider_icon(self, provider: str, icon_type: str, lang: str) -> tuple[bytes, str]:
|
def get_provider_icon(self, provider: str, icon_type: str, lang: str) -> tuple[bytes, str]:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@ -51,7 +51,7 @@ class ApiModeration(Moderation):
|
|||||||
params = ModerationInputParams(app_id=self.app_id, inputs=inputs, query=query)
|
params = ModerationInputParams(app_id=self.app_id, inputs=inputs, query=query)
|
||||||
|
|
||||||
result = self._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_INPUT, params.model_dump())
|
result = self._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_INPUT, params.model_dump())
|
||||||
return ModerationInputsResult(**result)
|
return ModerationInputsResult.model_validate(result)
|
||||||
|
|
||||||
return ModerationInputsResult(
|
return ModerationInputsResult(
|
||||||
flagged=flagged, action=ModerationAction.DIRECT_OUTPUT, preset_response=preset_response
|
flagged=flagged, action=ModerationAction.DIRECT_OUTPUT, preset_response=preset_response
|
||||||
@ -67,7 +67,7 @@ class ApiModeration(Moderation):
|
|||||||
params = ModerationOutputParams(app_id=self.app_id, text=text)
|
params = ModerationOutputParams(app_id=self.app_id, text=text)
|
||||||
|
|
||||||
result = self._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_OUTPUT, params.model_dump())
|
result = self._get_config_by_requestor(APIBasedExtensionPoint.APP_MODERATION_OUTPUT, params.model_dump())
|
||||||
return ModerationOutputsResult(**result)
|
return ModerationOutputsResult.model_validate(result)
|
||||||
|
|
||||||
return ModerationOutputsResult(
|
return ModerationOutputsResult(
|
||||||
flagged=flagged, action=ModerationAction.DIRECT_OUTPUT, preset_response=preset_response
|
flagged=flagged, action=ModerationAction.DIRECT_OUTPUT, preset_response=preset_response
|
||||||
|
|||||||
@ -213,9 +213,9 @@ class ArizePhoenixDataTrace(BaseTraceInstance):
|
|||||||
node_metadata.update(json.loads(node_execution.execution_metadata))
|
node_metadata.update(json.loads(node_execution.execution_metadata))
|
||||||
|
|
||||||
# Determine the correct span kind based on node type
|
# Determine the correct span kind based on node type
|
||||||
span_kind = OpenInferenceSpanKindValues.CHAIN.value
|
span_kind = OpenInferenceSpanKindValues.CHAIN
|
||||||
if node_execution.node_type == "llm":
|
if node_execution.node_type == "llm":
|
||||||
span_kind = OpenInferenceSpanKindValues.LLM.value
|
span_kind = OpenInferenceSpanKindValues.LLM
|
||||||
provider = process_data.get("model_provider")
|
provider = process_data.get("model_provider")
|
||||||
model = process_data.get("model_name")
|
model = process_data.get("model_name")
|
||||||
if provider:
|
if provider:
|
||||||
@ -230,18 +230,18 @@ class ArizePhoenixDataTrace(BaseTraceInstance):
|
|||||||
node_metadata["prompt_tokens"] = usage_data.get("prompt_tokens", 0)
|
node_metadata["prompt_tokens"] = usage_data.get("prompt_tokens", 0)
|
||||||
node_metadata["completion_tokens"] = usage_data.get("completion_tokens", 0)
|
node_metadata["completion_tokens"] = usage_data.get("completion_tokens", 0)
|
||||||
elif node_execution.node_type == "dataset_retrieval":
|
elif node_execution.node_type == "dataset_retrieval":
|
||||||
span_kind = OpenInferenceSpanKindValues.RETRIEVER.value
|
span_kind = OpenInferenceSpanKindValues.RETRIEVER
|
||||||
elif node_execution.node_type == "tool":
|
elif node_execution.node_type == "tool":
|
||||||
span_kind = OpenInferenceSpanKindValues.TOOL.value
|
span_kind = OpenInferenceSpanKindValues.TOOL
|
||||||
else:
|
else:
|
||||||
span_kind = OpenInferenceSpanKindValues.CHAIN.value
|
span_kind = OpenInferenceSpanKindValues.CHAIN
|
||||||
|
|
||||||
node_span = self.tracer.start_span(
|
node_span = self.tracer.start_span(
|
||||||
name=node_execution.node_type,
|
name=node_execution.node_type,
|
||||||
attributes={
|
attributes={
|
||||||
SpanAttributes.INPUT_VALUE: node_execution.inputs or "{}",
|
SpanAttributes.INPUT_VALUE: node_execution.inputs or "{}",
|
||||||
SpanAttributes.OUTPUT_VALUE: node_execution.outputs or "{}",
|
SpanAttributes.OUTPUT_VALUE: node_execution.outputs or "{}",
|
||||||
SpanAttributes.OPENINFERENCE_SPAN_KIND: span_kind,
|
SpanAttributes.OPENINFERENCE_SPAN_KIND: span_kind.value,
|
||||||
SpanAttributes.METADATA: json.dumps(node_metadata, ensure_ascii=False),
|
SpanAttributes.METADATA: json.dumps(node_metadata, ensure_ascii=False),
|
||||||
SpanAttributes.SESSION_ID: trace_info.conversation_id or "",
|
SpanAttributes.SESSION_ID: trace_info.conversation_id or "",
|
||||||
},
|
},
|
||||||
|
|||||||
@ -73,7 +73,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
if trace_info.message_id:
|
if trace_info.message_id:
|
||||||
trace_id = trace_info.trace_id or trace_info.message_id
|
trace_id = trace_info.trace_id or trace_info.message_id
|
||||||
name = TraceTaskName.MESSAGE_TRACE.value
|
name = TraceTaskName.MESSAGE_TRACE
|
||||||
trace_data = LangfuseTrace(
|
trace_data = LangfuseTrace(
|
||||||
id=trace_id,
|
id=trace_id,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
@ -88,7 +88,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
self.add_trace(langfuse_trace_data=trace_data)
|
self.add_trace(langfuse_trace_data=trace_data)
|
||||||
workflow_span_data = LangfuseSpan(
|
workflow_span_data = LangfuseSpan(
|
||||||
id=trace_info.workflow_run_id,
|
id=trace_info.workflow_run_id,
|
||||||
name=TraceTaskName.WORKFLOW_TRACE.value,
|
name=TraceTaskName.WORKFLOW_TRACE,
|
||||||
input=dict(trace_info.workflow_run_inputs),
|
input=dict(trace_info.workflow_run_inputs),
|
||||||
output=dict(trace_info.workflow_run_outputs),
|
output=dict(trace_info.workflow_run_outputs),
|
||||||
trace_id=trace_id,
|
trace_id=trace_id,
|
||||||
@ -103,7 +103,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
trace_data = LangfuseTrace(
|
trace_data = LangfuseTrace(
|
||||||
id=trace_id,
|
id=trace_id,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
name=TraceTaskName.WORKFLOW_TRACE.value,
|
name=TraceTaskName.WORKFLOW_TRACE,
|
||||||
input=dict(trace_info.workflow_run_inputs),
|
input=dict(trace_info.workflow_run_inputs),
|
||||||
output=dict(trace_info.workflow_run_outputs),
|
output=dict(trace_info.workflow_run_outputs),
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
@ -253,7 +253,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
trace_data = LangfuseTrace(
|
trace_data = LangfuseTrace(
|
||||||
id=trace_id,
|
id=trace_id,
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
name=TraceTaskName.MESSAGE_TRACE.value,
|
name=TraceTaskName.MESSAGE_TRACE,
|
||||||
input={
|
input={
|
||||||
"message": trace_info.inputs,
|
"message": trace_info.inputs,
|
||||||
"files": file_list,
|
"files": file_list,
|
||||||
@ -303,7 +303,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
if trace_info.message_data is None:
|
if trace_info.message_data is None:
|
||||||
return
|
return
|
||||||
span_data = LangfuseSpan(
|
span_data = LangfuseSpan(
|
||||||
name=TraceTaskName.MODERATION_TRACE.value,
|
name=TraceTaskName.MODERATION_TRACE,
|
||||||
input=trace_info.inputs,
|
input=trace_info.inputs,
|
||||||
output={
|
output={
|
||||||
"action": trace_info.action,
|
"action": trace_info.action,
|
||||||
@ -331,7 +331,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
)
|
)
|
||||||
|
|
||||||
generation_data = LangfuseGeneration(
|
generation_data = LangfuseGeneration(
|
||||||
name=TraceTaskName.SUGGESTED_QUESTION_TRACE.value,
|
name=TraceTaskName.SUGGESTED_QUESTION_TRACE,
|
||||||
input=trace_info.inputs,
|
input=trace_info.inputs,
|
||||||
output=str(trace_info.suggested_question),
|
output=str(trace_info.suggested_question),
|
||||||
trace_id=trace_info.trace_id or trace_info.message_id,
|
trace_id=trace_info.trace_id or trace_info.message_id,
|
||||||
@ -349,7 +349,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
if trace_info.message_data is None:
|
if trace_info.message_data is None:
|
||||||
return
|
return
|
||||||
dataset_retrieval_span_data = LangfuseSpan(
|
dataset_retrieval_span_data = LangfuseSpan(
|
||||||
name=TraceTaskName.DATASET_RETRIEVAL_TRACE.value,
|
name=TraceTaskName.DATASET_RETRIEVAL_TRACE,
|
||||||
input=trace_info.inputs,
|
input=trace_info.inputs,
|
||||||
output={"documents": trace_info.documents},
|
output={"documents": trace_info.documents},
|
||||||
trace_id=trace_info.trace_id or trace_info.message_id,
|
trace_id=trace_info.trace_id or trace_info.message_id,
|
||||||
@ -377,7 +377,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
def generate_name_trace(self, trace_info: GenerateNameTraceInfo):
|
def generate_name_trace(self, trace_info: GenerateNameTraceInfo):
|
||||||
name_generation_trace_data = LangfuseTrace(
|
name_generation_trace_data = LangfuseTrace(
|
||||||
name=TraceTaskName.GENERATE_NAME_TRACE.value,
|
name=TraceTaskName.GENERATE_NAME_TRACE,
|
||||||
input=trace_info.inputs,
|
input=trace_info.inputs,
|
||||||
output=trace_info.outputs,
|
output=trace_info.outputs,
|
||||||
user_id=trace_info.tenant_id,
|
user_id=trace_info.tenant_id,
|
||||||
@ -388,7 +388,7 @@ class LangFuseDataTrace(BaseTraceInstance):
|
|||||||
self.add_trace(langfuse_trace_data=name_generation_trace_data)
|
self.add_trace(langfuse_trace_data=name_generation_trace_data)
|
||||||
|
|
||||||
name_generation_span_data = LangfuseSpan(
|
name_generation_span_data = LangfuseSpan(
|
||||||
name=TraceTaskName.GENERATE_NAME_TRACE.value,
|
name=TraceTaskName.GENERATE_NAME_TRACE,
|
||||||
input=trace_info.inputs,
|
input=trace_info.inputs,
|
||||||
output=trace_info.outputs,
|
output=trace_info.outputs,
|
||||||
trace_id=trace_info.conversation_id,
|
trace_id=trace_info.conversation_id,
|
||||||
|
|||||||
@ -81,7 +81,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
if trace_info.message_id:
|
if trace_info.message_id:
|
||||||
message_run = LangSmithRunModel(
|
message_run = LangSmithRunModel(
|
||||||
id=trace_info.message_id,
|
id=trace_info.message_id,
|
||||||
name=TraceTaskName.MESSAGE_TRACE.value,
|
name=TraceTaskName.MESSAGE_TRACE,
|
||||||
inputs=dict(trace_info.workflow_run_inputs),
|
inputs=dict(trace_info.workflow_run_inputs),
|
||||||
outputs=dict(trace_info.workflow_run_outputs),
|
outputs=dict(trace_info.workflow_run_outputs),
|
||||||
run_type=LangSmithRunType.chain,
|
run_type=LangSmithRunType.chain,
|
||||||
@ -110,7 +110,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
file_list=trace_info.file_list,
|
file_list=trace_info.file_list,
|
||||||
total_tokens=trace_info.total_tokens,
|
total_tokens=trace_info.total_tokens,
|
||||||
id=trace_info.workflow_run_id,
|
id=trace_info.workflow_run_id,
|
||||||
name=TraceTaskName.WORKFLOW_TRACE.value,
|
name=TraceTaskName.WORKFLOW_TRACE,
|
||||||
inputs=dict(trace_info.workflow_run_inputs),
|
inputs=dict(trace_info.workflow_run_inputs),
|
||||||
run_type=LangSmithRunType.tool,
|
run_type=LangSmithRunType.tool,
|
||||||
start_time=trace_info.workflow_data.created_at,
|
start_time=trace_info.workflow_data.created_at,
|
||||||
@ -271,7 +271,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
output_tokens=trace_info.answer_tokens,
|
output_tokens=trace_info.answer_tokens,
|
||||||
total_tokens=trace_info.total_tokens,
|
total_tokens=trace_info.total_tokens,
|
||||||
id=message_id,
|
id=message_id,
|
||||||
name=TraceTaskName.MESSAGE_TRACE.value,
|
name=TraceTaskName.MESSAGE_TRACE,
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
run_type=LangSmithRunType.chain,
|
run_type=LangSmithRunType.chain,
|
||||||
start_time=trace_info.start_time,
|
start_time=trace_info.start_time,
|
||||||
@ -327,7 +327,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
if trace_info.message_data is None:
|
if trace_info.message_data is None:
|
||||||
return
|
return
|
||||||
langsmith_run = LangSmithRunModel(
|
langsmith_run = LangSmithRunModel(
|
||||||
name=TraceTaskName.MODERATION_TRACE.value,
|
name=TraceTaskName.MODERATION_TRACE,
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs={
|
outputs={
|
||||||
"action": trace_info.action,
|
"action": trace_info.action,
|
||||||
@ -362,7 +362,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
if message_data is None:
|
if message_data is None:
|
||||||
return
|
return
|
||||||
suggested_question_run = LangSmithRunModel(
|
suggested_question_run = LangSmithRunModel(
|
||||||
name=TraceTaskName.SUGGESTED_QUESTION_TRACE.value,
|
name=TraceTaskName.SUGGESTED_QUESTION_TRACE,
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs=trace_info.suggested_question,
|
outputs=trace_info.suggested_question,
|
||||||
run_type=LangSmithRunType.tool,
|
run_type=LangSmithRunType.tool,
|
||||||
@ -391,7 +391,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
if trace_info.message_data is None:
|
if trace_info.message_data is None:
|
||||||
return
|
return
|
||||||
dataset_retrieval_run = LangSmithRunModel(
|
dataset_retrieval_run = LangSmithRunModel(
|
||||||
name=TraceTaskName.DATASET_RETRIEVAL_TRACE.value,
|
name=TraceTaskName.DATASET_RETRIEVAL_TRACE,
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs={"documents": trace_info.documents},
|
outputs={"documents": trace_info.documents},
|
||||||
run_type=LangSmithRunType.retriever,
|
run_type=LangSmithRunType.retriever,
|
||||||
@ -447,7 +447,7 @@ class LangSmithDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
def generate_name_trace(self, trace_info: GenerateNameTraceInfo):
|
def generate_name_trace(self, trace_info: GenerateNameTraceInfo):
|
||||||
name_run = LangSmithRunModel(
|
name_run = LangSmithRunModel(
|
||||||
name=TraceTaskName.GENERATE_NAME_TRACE.value,
|
name=TraceTaskName.GENERATE_NAME_TRACE,
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs=trace_info.outputs,
|
outputs=trace_info.outputs,
|
||||||
run_type=LangSmithRunType.tool,
|
run_type=LangSmithRunType.tool,
|
||||||
|
|||||||
@ -108,7 +108,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
trace_data = {
|
trace_data = {
|
||||||
"id": opik_trace_id,
|
"id": opik_trace_id,
|
||||||
"name": TraceTaskName.MESSAGE_TRACE.value,
|
"name": TraceTaskName.MESSAGE_TRACE,
|
||||||
"start_time": trace_info.start_time,
|
"start_time": trace_info.start_time,
|
||||||
"end_time": trace_info.end_time,
|
"end_time": trace_info.end_time,
|
||||||
"metadata": workflow_metadata,
|
"metadata": workflow_metadata,
|
||||||
@ -125,7 +125,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
"id": root_span_id,
|
"id": root_span_id,
|
||||||
"parent_span_id": None,
|
"parent_span_id": None,
|
||||||
"trace_id": opik_trace_id,
|
"trace_id": opik_trace_id,
|
||||||
"name": TraceTaskName.WORKFLOW_TRACE.value,
|
"name": TraceTaskName.WORKFLOW_TRACE,
|
||||||
"input": wrap_dict("input", trace_info.workflow_run_inputs),
|
"input": wrap_dict("input", trace_info.workflow_run_inputs),
|
||||||
"output": wrap_dict("output", trace_info.workflow_run_outputs),
|
"output": wrap_dict("output", trace_info.workflow_run_outputs),
|
||||||
"start_time": trace_info.start_time,
|
"start_time": trace_info.start_time,
|
||||||
@ -138,7 +138,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
else:
|
else:
|
||||||
trace_data = {
|
trace_data = {
|
||||||
"id": opik_trace_id,
|
"id": opik_trace_id,
|
||||||
"name": TraceTaskName.MESSAGE_TRACE.value,
|
"name": TraceTaskName.MESSAGE_TRACE,
|
||||||
"start_time": trace_info.start_time,
|
"start_time": trace_info.start_time,
|
||||||
"end_time": trace_info.end_time,
|
"end_time": trace_info.end_time,
|
||||||
"metadata": workflow_metadata,
|
"metadata": workflow_metadata,
|
||||||
@ -290,7 +290,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
trace_data = {
|
trace_data = {
|
||||||
"id": prepare_opik_uuid(trace_info.start_time, dify_trace_id),
|
"id": prepare_opik_uuid(trace_info.start_time, dify_trace_id),
|
||||||
"name": TraceTaskName.MESSAGE_TRACE.value,
|
"name": TraceTaskName.MESSAGE_TRACE,
|
||||||
"start_time": trace_info.start_time,
|
"start_time": trace_info.start_time,
|
||||||
"end_time": trace_info.end_time,
|
"end_time": trace_info.end_time,
|
||||||
"metadata": wrap_metadata(metadata),
|
"metadata": wrap_metadata(metadata),
|
||||||
@ -329,7 +329,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
span_data = {
|
span_data = {
|
||||||
"trace_id": prepare_opik_uuid(start_time, trace_info.trace_id or trace_info.message_id),
|
"trace_id": prepare_opik_uuid(start_time, trace_info.trace_id or trace_info.message_id),
|
||||||
"name": TraceTaskName.MODERATION_TRACE.value,
|
"name": TraceTaskName.MODERATION_TRACE,
|
||||||
"type": "tool",
|
"type": "tool",
|
||||||
"start_time": start_time,
|
"start_time": start_time,
|
||||||
"end_time": trace_info.end_time or trace_info.message_data.updated_at,
|
"end_time": trace_info.end_time or trace_info.message_data.updated_at,
|
||||||
@ -355,7 +355,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
span_data = {
|
span_data = {
|
||||||
"trace_id": prepare_opik_uuid(start_time, trace_info.trace_id or trace_info.message_id),
|
"trace_id": prepare_opik_uuid(start_time, trace_info.trace_id or trace_info.message_id),
|
||||||
"name": TraceTaskName.SUGGESTED_QUESTION_TRACE.value,
|
"name": TraceTaskName.SUGGESTED_QUESTION_TRACE,
|
||||||
"type": "tool",
|
"type": "tool",
|
||||||
"start_time": start_time,
|
"start_time": start_time,
|
||||||
"end_time": trace_info.end_time or message_data.updated_at,
|
"end_time": trace_info.end_time or message_data.updated_at,
|
||||||
@ -375,7 +375,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
span_data = {
|
span_data = {
|
||||||
"trace_id": prepare_opik_uuid(start_time, trace_info.trace_id or trace_info.message_id),
|
"trace_id": prepare_opik_uuid(start_time, trace_info.trace_id or trace_info.message_id),
|
||||||
"name": TraceTaskName.DATASET_RETRIEVAL_TRACE.value,
|
"name": TraceTaskName.DATASET_RETRIEVAL_TRACE,
|
||||||
"type": "tool",
|
"type": "tool",
|
||||||
"start_time": start_time,
|
"start_time": start_time,
|
||||||
"end_time": trace_info.end_time or trace_info.message_data.updated_at,
|
"end_time": trace_info.end_time or trace_info.message_data.updated_at,
|
||||||
@ -405,7 +405,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
def generate_name_trace(self, trace_info: GenerateNameTraceInfo):
|
def generate_name_trace(self, trace_info: GenerateNameTraceInfo):
|
||||||
trace_data = {
|
trace_data = {
|
||||||
"id": prepare_opik_uuid(trace_info.start_time, trace_info.trace_id or trace_info.message_id),
|
"id": prepare_opik_uuid(trace_info.start_time, trace_info.trace_id or trace_info.message_id),
|
||||||
"name": TraceTaskName.GENERATE_NAME_TRACE.value,
|
"name": TraceTaskName.GENERATE_NAME_TRACE,
|
||||||
"start_time": trace_info.start_time,
|
"start_time": trace_info.start_time,
|
||||||
"end_time": trace_info.end_time,
|
"end_time": trace_info.end_time,
|
||||||
"metadata": wrap_metadata(trace_info.metadata),
|
"metadata": wrap_metadata(trace_info.metadata),
|
||||||
@ -420,7 +420,7 @@ class OpikDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
span_data = {
|
span_data = {
|
||||||
"trace_id": trace.id,
|
"trace_id": trace.id,
|
||||||
"name": TraceTaskName.GENERATE_NAME_TRACE.value,
|
"name": TraceTaskName.GENERATE_NAME_TRACE,
|
||||||
"start_time": trace_info.start_time,
|
"start_time": trace_info.start_time,
|
||||||
"end_time": trace_info.end_time,
|
"end_time": trace_info.end_time,
|
||||||
"metadata": wrap_metadata(trace_info.metadata),
|
"metadata": wrap_metadata(trace_info.metadata),
|
||||||
|
|||||||
@ -104,7 +104,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
message_run = WeaveTraceModel(
|
message_run = WeaveTraceModel(
|
||||||
id=trace_info.message_id,
|
id=trace_info.message_id,
|
||||||
op=str(TraceTaskName.MESSAGE_TRACE.value),
|
op=str(TraceTaskName.MESSAGE_TRACE),
|
||||||
inputs=dict(trace_info.workflow_run_inputs),
|
inputs=dict(trace_info.workflow_run_inputs),
|
||||||
outputs=dict(trace_info.workflow_run_outputs),
|
outputs=dict(trace_info.workflow_run_outputs),
|
||||||
total_tokens=trace_info.total_tokens,
|
total_tokens=trace_info.total_tokens,
|
||||||
@ -126,7 +126,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
file_list=trace_info.file_list,
|
file_list=trace_info.file_list,
|
||||||
total_tokens=trace_info.total_tokens,
|
total_tokens=trace_info.total_tokens,
|
||||||
id=trace_info.workflow_run_id,
|
id=trace_info.workflow_run_id,
|
||||||
op=str(TraceTaskName.WORKFLOW_TRACE.value),
|
op=str(TraceTaskName.WORKFLOW_TRACE),
|
||||||
inputs=dict(trace_info.workflow_run_inputs),
|
inputs=dict(trace_info.workflow_run_inputs),
|
||||||
outputs=dict(trace_info.workflow_run_outputs),
|
outputs=dict(trace_info.workflow_run_outputs),
|
||||||
attributes=workflow_attributes,
|
attributes=workflow_attributes,
|
||||||
@ -253,7 +253,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
message_run = WeaveTraceModel(
|
message_run = WeaveTraceModel(
|
||||||
id=trace_id,
|
id=trace_id,
|
||||||
op=str(TraceTaskName.MESSAGE_TRACE.value),
|
op=str(TraceTaskName.MESSAGE_TRACE),
|
||||||
input_tokens=trace_info.message_tokens,
|
input_tokens=trace_info.message_tokens,
|
||||||
output_tokens=trace_info.answer_tokens,
|
output_tokens=trace_info.answer_tokens,
|
||||||
total_tokens=trace_info.total_tokens,
|
total_tokens=trace_info.total_tokens,
|
||||||
@ -300,7 +300,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
moderation_run = WeaveTraceModel(
|
moderation_run = WeaveTraceModel(
|
||||||
id=str(uuid.uuid4()),
|
id=str(uuid.uuid4()),
|
||||||
op=str(TraceTaskName.MODERATION_TRACE.value),
|
op=str(TraceTaskName.MODERATION_TRACE),
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs={
|
outputs={
|
||||||
"action": trace_info.action,
|
"action": trace_info.action,
|
||||||
@ -330,7 +330,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
suggested_question_run = WeaveTraceModel(
|
suggested_question_run = WeaveTraceModel(
|
||||||
id=str(uuid.uuid4()),
|
id=str(uuid.uuid4()),
|
||||||
op=str(TraceTaskName.SUGGESTED_QUESTION_TRACE.value),
|
op=str(TraceTaskName.SUGGESTED_QUESTION_TRACE),
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs=trace_info.suggested_question,
|
outputs=trace_info.suggested_question,
|
||||||
attributes=attributes,
|
attributes=attributes,
|
||||||
@ -355,7 +355,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
dataset_retrieval_run = WeaveTraceModel(
|
dataset_retrieval_run = WeaveTraceModel(
|
||||||
id=str(uuid.uuid4()),
|
id=str(uuid.uuid4()),
|
||||||
op=str(TraceTaskName.DATASET_RETRIEVAL_TRACE.value),
|
op=str(TraceTaskName.DATASET_RETRIEVAL_TRACE),
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs={"documents": trace_info.documents},
|
outputs={"documents": trace_info.documents},
|
||||||
attributes=attributes,
|
attributes=attributes,
|
||||||
@ -397,7 +397,7 @@ class WeaveDataTrace(BaseTraceInstance):
|
|||||||
|
|
||||||
name_run = WeaveTraceModel(
|
name_run = WeaveTraceModel(
|
||||||
id=str(uuid.uuid4()),
|
id=str(uuid.uuid4()),
|
||||||
op=str(TraceTaskName.GENERATE_NAME_TRACE.value),
|
op=str(TraceTaskName.GENERATE_NAME_TRACE),
|
||||||
inputs=trace_info.inputs,
|
inputs=trace_info.inputs,
|
||||||
outputs=trace_info.outputs,
|
outputs=trace_info.outputs,
|
||||||
attributes=attributes,
|
attributes=attributes,
|
||||||
|
|||||||
@ -52,7 +52,7 @@ class PluginNodeBackwardsInvocation(BaseBackwardsInvocation):
|
|||||||
instruction=instruction, # instruct with variables are not supported
|
instruction=instruction, # instruct with variables are not supported
|
||||||
)
|
)
|
||||||
node_data_dict = node_data.model_dump()
|
node_data_dict = node_data.model_dump()
|
||||||
node_data_dict["type"] = NodeType.PARAMETER_EXTRACTOR.value
|
node_data_dict["type"] = NodeType.PARAMETER_EXTRACTOR
|
||||||
execution = workflow_service.run_free_workflow_node(
|
execution = workflow_service.run_free_workflow_node(
|
||||||
node_data_dict,
|
node_data_dict,
|
||||||
tenant_id=tenant_id,
|
tenant_id=tenant_id,
|
||||||
|
|||||||
@ -83,16 +83,16 @@ class RequestInvokeLLM(BaseRequestInvokeModel):
|
|||||||
raise ValueError("prompt_messages must be a list")
|
raise ValueError("prompt_messages must be a list")
|
||||||
|
|
||||||
for i in range(len(v)):
|
for i in range(len(v)):
|
||||||
if v[i]["role"] == PromptMessageRole.USER.value:
|
if v[i]["role"] == PromptMessageRole.USER:
|
||||||
v[i] = UserPromptMessage(**v[i])
|
v[i] = UserPromptMessage.model_validate(v[i])
|
||||||
elif v[i]["role"] == PromptMessageRole.ASSISTANT.value:
|
elif v[i]["role"] == PromptMessageRole.ASSISTANT:
|
||||||
v[i] = AssistantPromptMessage(**v[i])
|
v[i] = AssistantPromptMessage.model_validate(v[i])
|
||||||
elif v[i]["role"] == PromptMessageRole.SYSTEM.value:
|
elif v[i]["role"] == PromptMessageRole.SYSTEM:
|
||||||
v[i] = SystemPromptMessage(**v[i])
|
v[i] = SystemPromptMessage.model_validate(v[i])
|
||||||
elif v[i]["role"] == PromptMessageRole.TOOL.value:
|
elif v[i]["role"] == PromptMessageRole.TOOL:
|
||||||
v[i] = ToolPromptMessage(**v[i])
|
v[i] = ToolPromptMessage.model_validate(v[i])
|
||||||
else:
|
else:
|
||||||
v[i] = PromptMessage(**v[i])
|
v[i] = PromptMessage.model_validate(v[i])
|
||||||
|
|
||||||
return v
|
return v
|
||||||
|
|
||||||
|
|||||||
@ -2,11 +2,10 @@ import inspect
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
from collections.abc import Callable, Generator
|
from collections.abc import Callable, Generator
|
||||||
from typing import TypeVar
|
from typing import Any, TypeVar
|
||||||
|
|
||||||
import requests
|
import httpx
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
from requests.exceptions import HTTPError
|
|
||||||
from yarl import URL
|
from yarl import URL
|
||||||
|
|
||||||
from configs import dify_config
|
from configs import dify_config
|
||||||
@ -47,29 +46,56 @@ class BasePluginClient:
|
|||||||
data: bytes | dict | str | None = None,
|
data: bytes | dict | str | None = None,
|
||||||
params: dict | None = None,
|
params: dict | None = None,
|
||||||
files: dict | None = None,
|
files: dict | None = None,
|
||||||
stream: bool = False,
|
) -> httpx.Response:
|
||||||
) -> requests.Response:
|
|
||||||
"""
|
"""
|
||||||
Make a request to the plugin daemon inner API.
|
Make a request to the plugin daemon inner API.
|
||||||
"""
|
"""
|
||||||
url = plugin_daemon_inner_api_baseurl / path
|
url, headers, prepared_data, params, files = self._prepare_request(path, headers, data, params, files)
|
||||||
headers = headers or {}
|
|
||||||
headers["X-Api-Key"] = dify_config.PLUGIN_DAEMON_KEY
|
|
||||||
headers["Accept-Encoding"] = "gzip, deflate, br"
|
|
||||||
|
|
||||||
if headers.get("Content-Type") == "application/json" and isinstance(data, dict):
|
request_kwargs: dict[str, Any] = {
|
||||||
data = json.dumps(data)
|
"method": method,
|
||||||
|
"url": url,
|
||||||
|
"headers": headers,
|
||||||
|
"params": params,
|
||||||
|
"files": files,
|
||||||
|
}
|
||||||
|
if isinstance(prepared_data, dict):
|
||||||
|
request_kwargs["data"] = prepared_data
|
||||||
|
elif prepared_data is not None:
|
||||||
|
request_kwargs["content"] = prepared_data
|
||||||
|
|
||||||
try:
|
try:
|
||||||
response = requests.request(
|
response = httpx.request(**request_kwargs)
|
||||||
method=method, url=str(url), headers=headers, data=data, params=params, stream=stream, files=files
|
except httpx.RequestError:
|
||||||
)
|
|
||||||
except requests.ConnectionError:
|
|
||||||
logger.exception("Request to Plugin Daemon Service failed")
|
logger.exception("Request to Plugin Daemon Service failed")
|
||||||
raise PluginDaemonInnerError(code=-500, message="Request to Plugin Daemon Service failed")
|
raise PluginDaemonInnerError(code=-500, message="Request to Plugin Daemon Service failed")
|
||||||
|
|
||||||
return response
|
return response
|
||||||
|
|
||||||
|
def _prepare_request(
|
||||||
|
self,
|
||||||
|
path: str,
|
||||||
|
headers: dict | None,
|
||||||
|
data: bytes | dict | str | None,
|
||||||
|
params: dict | None,
|
||||||
|
files: dict | None,
|
||||||
|
) -> tuple[str, dict, bytes | dict | str | None, dict | None, dict | None]:
|
||||||
|
url = plugin_daemon_inner_api_baseurl / path
|
||||||
|
prepared_headers = dict(headers or {})
|
||||||
|
prepared_headers["X-Api-Key"] = dify_config.PLUGIN_DAEMON_KEY
|
||||||
|
prepared_headers.setdefault("Accept-Encoding", "gzip, deflate, br")
|
||||||
|
|
||||||
|
prepared_data: bytes | dict | str | None = (
|
||||||
|
data if isinstance(data, (bytes, str, dict)) or data is None else None
|
||||||
|
)
|
||||||
|
if isinstance(data, dict):
|
||||||
|
if prepared_headers.get("Content-Type") == "application/json":
|
||||||
|
prepared_data = json.dumps(data)
|
||||||
|
else:
|
||||||
|
prepared_data = data
|
||||||
|
|
||||||
|
return str(url), prepared_headers, prepared_data, params, files
|
||||||
|
|
||||||
def _stream_request(
|
def _stream_request(
|
||||||
self,
|
self,
|
||||||
method: str,
|
method: str,
|
||||||
@ -78,23 +104,44 @@ class BasePluginClient:
|
|||||||
headers: dict | None = None,
|
headers: dict | None = None,
|
||||||
data: bytes | dict | None = None,
|
data: bytes | dict | None = None,
|
||||||
files: dict | None = None,
|
files: dict | None = None,
|
||||||
) -> Generator[bytes, None, None]:
|
) -> Generator[str, None, None]:
|
||||||
"""
|
"""
|
||||||
Make a stream request to the plugin daemon inner API
|
Make a stream request to the plugin daemon inner API
|
||||||
"""
|
"""
|
||||||
response = self._request(method, path, headers, data, params, files, stream=True)
|
url, headers, prepared_data, params, files = self._prepare_request(path, headers, data, params, files)
|
||||||
for line in response.iter_lines(chunk_size=1024 * 8):
|
|
||||||
line = line.decode("utf-8").strip()
|
stream_kwargs: dict[str, Any] = {
|
||||||
if line.startswith("data:"):
|
"method": method,
|
||||||
line = line[5:].strip()
|
"url": url,
|
||||||
if line:
|
"headers": headers,
|
||||||
yield line
|
"params": params,
|
||||||
|
"files": files,
|
||||||
|
}
|
||||||
|
if isinstance(prepared_data, dict):
|
||||||
|
stream_kwargs["data"] = prepared_data
|
||||||
|
elif prepared_data is not None:
|
||||||
|
stream_kwargs["content"] = prepared_data
|
||||||
|
|
||||||
|
try:
|
||||||
|
with httpx.stream(**stream_kwargs) as response:
|
||||||
|
for raw_line in response.iter_lines():
|
||||||
|
if raw_line is None:
|
||||||
|
continue
|
||||||
|
line = raw_line.decode("utf-8") if isinstance(raw_line, bytes) else raw_line
|
||||||
|
line = line.strip()
|
||||||
|
if line.startswith("data:"):
|
||||||
|
line = line[5:].strip()
|
||||||
|
if line:
|
||||||
|
yield line
|
||||||
|
except httpx.RequestError:
|
||||||
|
logger.exception("Stream request to Plugin Daemon Service failed")
|
||||||
|
raise PluginDaemonInnerError(code=-500, message="Request to Plugin Daemon Service failed")
|
||||||
|
|
||||||
def _stream_request_with_model(
|
def _stream_request_with_model(
|
||||||
self,
|
self,
|
||||||
method: str,
|
method: str,
|
||||||
path: str,
|
path: str,
|
||||||
type: type[T],
|
type_: type[T],
|
||||||
headers: dict | None = None,
|
headers: dict | None = None,
|
||||||
data: bytes | dict | None = None,
|
data: bytes | dict | None = None,
|
||||||
params: dict | None = None,
|
params: dict | None = None,
|
||||||
@ -104,13 +151,13 @@ class BasePluginClient:
|
|||||||
Make a stream request to the plugin daemon inner API and yield the response as a model.
|
Make a stream request to the plugin daemon inner API and yield the response as a model.
|
||||||
"""
|
"""
|
||||||
for line in self._stream_request(method, path, params, headers, data, files):
|
for line in self._stream_request(method, path, params, headers, data, files):
|
||||||
yield type(**json.loads(line)) # type: ignore
|
yield type_(**json.loads(line)) # type: ignore
|
||||||
|
|
||||||
def _request_with_model(
|
def _request_with_model(
|
||||||
self,
|
self,
|
||||||
method: str,
|
method: str,
|
||||||
path: str,
|
path: str,
|
||||||
type: type[T],
|
type_: type[T],
|
||||||
headers: dict | None = None,
|
headers: dict | None = None,
|
||||||
data: bytes | None = None,
|
data: bytes | None = None,
|
||||||
params: dict | None = None,
|
params: dict | None = None,
|
||||||
@ -120,13 +167,13 @@ class BasePluginClient:
|
|||||||
Make a request to the plugin daemon inner API and return the response as a model.
|
Make a request to the plugin daemon inner API and return the response as a model.
|
||||||
"""
|
"""
|
||||||
response = self._request(method, path, headers, data, params, files)
|
response = self._request(method, path, headers, data, params, files)
|
||||||
return type(**response.json()) # type: ignore
|
return type_(**response.json()) # type: ignore
|
||||||
|
|
||||||
def _request_with_plugin_daemon_response(
|
def _request_with_plugin_daemon_response(
|
||||||
self,
|
self,
|
||||||
method: str,
|
method: str,
|
||||||
path: str,
|
path: str,
|
||||||
type: type[T],
|
type_: type[T],
|
||||||
headers: dict | None = None,
|
headers: dict | None = None,
|
||||||
data: bytes | dict | None = None,
|
data: bytes | dict | None = None,
|
||||||
params: dict | None = None,
|
params: dict | None = None,
|
||||||
@ -139,23 +186,23 @@ class BasePluginClient:
|
|||||||
try:
|
try:
|
||||||
response = self._request(method, path, headers, data, params, files)
|
response = self._request(method, path, headers, data, params, files)
|
||||||
response.raise_for_status()
|
response.raise_for_status()
|
||||||
except HTTPError as e:
|
except httpx.HTTPStatusError as e:
|
||||||
msg = f"Failed to request plugin daemon, status: {e.response.status_code}, url: {path}"
|
logger.exception("Failed to request plugin daemon, status: %s, url: %s", e.response.status_code, path)
|
||||||
logger.exception(msg)
|
|
||||||
raise e
|
raise e
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
msg = f"Failed to request plugin daemon, url: {path}"
|
msg = f"Failed to request plugin daemon, url: {path}"
|
||||||
logger.exception(msg)
|
logger.exception("Failed to request plugin daemon, url: %s", path)
|
||||||
raise ValueError(msg) from e
|
raise ValueError(msg) from e
|
||||||
|
|
||||||
try:
|
try:
|
||||||
json_response = response.json()
|
json_response = response.json()
|
||||||
if transformer:
|
if transformer:
|
||||||
json_response = transformer(json_response)
|
json_response = transformer(json_response)
|
||||||
rep = PluginDaemonBasicResponse[type](**json_response) # type: ignore
|
# https://stackoverflow.com/questions/59634937/variable-foo-class-is-not-valid-as-type-but-why
|
||||||
|
rep = PluginDaemonBasicResponse[type_].model_validate(json_response) # type: ignore
|
||||||
except Exception:
|
except Exception:
|
||||||
msg = (
|
msg = (
|
||||||
f"Failed to parse response from plugin daemon to PluginDaemonBasicResponse [{str(type.__name__)}],"
|
f"Failed to parse response from plugin daemon to PluginDaemonBasicResponse [{str(type_.__name__)}],"
|
||||||
f" url: {path}"
|
f" url: {path}"
|
||||||
)
|
)
|
||||||
logger.exception(msg)
|
logger.exception(msg)
|
||||||
@ -163,7 +210,7 @@ class BasePluginClient:
|
|||||||
|
|
||||||
if rep.code != 0:
|
if rep.code != 0:
|
||||||
try:
|
try:
|
||||||
error = PluginDaemonError(**json.loads(rep.message))
|
error = PluginDaemonError.model_validate(json.loads(rep.message))
|
||||||
except Exception:
|
except Exception:
|
||||||
raise ValueError(f"{rep.message}, code: {rep.code}")
|
raise ValueError(f"{rep.message}, code: {rep.code}")
|
||||||
|
|
||||||
@ -178,7 +225,7 @@ class BasePluginClient:
|
|||||||
self,
|
self,
|
||||||
method: str,
|
method: str,
|
||||||
path: str,
|
path: str,
|
||||||
type: type[T],
|
type_: type[T],
|
||||||
headers: dict | None = None,
|
headers: dict | None = None,
|
||||||
data: bytes | dict | None = None,
|
data: bytes | dict | None = None,
|
||||||
params: dict | None = None,
|
params: dict | None = None,
|
||||||
@ -189,7 +236,7 @@ class BasePluginClient:
|
|||||||
"""
|
"""
|
||||||
for line in self._stream_request(method, path, params, headers, data, files):
|
for line in self._stream_request(method, path, params, headers, data, files):
|
||||||
try:
|
try:
|
||||||
rep = PluginDaemonBasicResponse[type].model_validate_json(line) # type: ignore
|
rep = PluginDaemonBasicResponse[type_].model_validate_json(line) # type: ignore
|
||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
# TODO modify this when line_data has code and message
|
# TODO modify this when line_data has code and message
|
||||||
try:
|
try:
|
||||||
@ -204,11 +251,11 @@ class BasePluginClient:
|
|||||||
if rep.code != 0:
|
if rep.code != 0:
|
||||||
if rep.code == -500:
|
if rep.code == -500:
|
||||||
try:
|
try:
|
||||||
error = PluginDaemonError(**json.loads(rep.message))
|
error = PluginDaemonError.model_validate(json.loads(rep.message))
|
||||||
except Exception:
|
except Exception:
|
||||||
raise PluginDaemonInnerError(code=rep.code, message=rep.message)
|
raise PluginDaemonInnerError(code=rep.code, message=rep.message)
|
||||||
|
|
||||||
logger.error("Error in stream reponse for plugin %s", rep.__dict__)
|
logger.error("Error in stream response for plugin %s", rep.__dict__)
|
||||||
self._handle_plugin_daemon_error(error.error_type, error.message)
|
self._handle_plugin_daemon_error(error.error_type, error.message)
|
||||||
raise ValueError(f"plugin daemon: {rep.message}, code: {rep.code}")
|
raise ValueError(f"plugin daemon: {rep.message}, code: {rep.code}")
|
||||||
if rep.data is None:
|
if rep.data is None:
|
||||||
|
|||||||
@ -46,7 +46,9 @@ class PluginDatasourceManager(BasePluginClient):
|
|||||||
params={"page": 1, "page_size": 256},
|
params={"page": 1, "page_size": 256},
|
||||||
transformer=transformer,
|
transformer=transformer,
|
||||||
)
|
)
|
||||||
local_file_datasource_provider = PluginDatasourceProviderEntity(**self._get_local_file_datasource_provider())
|
local_file_datasource_provider = PluginDatasourceProviderEntity.model_validate(
|
||||||
|
self._get_local_file_datasource_provider()
|
||||||
|
)
|
||||||
|
|
||||||
for provider in response:
|
for provider in response:
|
||||||
ToolTransformService.repack_provider(tenant_id=tenant_id, provider=provider)
|
ToolTransformService.repack_provider(tenant_id=tenant_id, provider=provider)
|
||||||
@ -104,7 +106,7 @@ class PluginDatasourceManager(BasePluginClient):
|
|||||||
Fetch datasource provider for the given tenant and plugin.
|
Fetch datasource provider for the given tenant and plugin.
|
||||||
"""
|
"""
|
||||||
if provider_id == "langgenius/file/file":
|
if provider_id == "langgenius/file/file":
|
||||||
return PluginDatasourceProviderEntity(**self._get_local_file_datasource_provider())
|
return PluginDatasourceProviderEntity.model_validate(self._get_local_file_datasource_provider())
|
||||||
|
|
||||||
tool_provider_id = DatasourceProviderID(provider_id)
|
tool_provider_id = DatasourceProviderID(provider_id)
|
||||||
|
|
||||||
|
|||||||
@ -162,7 +162,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/llm/invoke",
|
path=f"plugin/{tenant_id}/dispatch/llm/invoke",
|
||||||
type=LLMResultChunk,
|
type_=LLMResultChunk,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -208,7 +208,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/llm/num_tokens",
|
path=f"plugin/{tenant_id}/dispatch/llm/num_tokens",
|
||||||
type=PluginLLMNumTokensResponse,
|
type_=PluginLLMNumTokensResponse,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -250,7 +250,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/text_embedding/invoke",
|
path=f"plugin/{tenant_id}/dispatch/text_embedding/invoke",
|
||||||
type=TextEmbeddingResult,
|
type_=TextEmbeddingResult,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -291,7 +291,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/text_embedding/num_tokens",
|
path=f"plugin/{tenant_id}/dispatch/text_embedding/num_tokens",
|
||||||
type=PluginTextEmbeddingNumTokensResponse,
|
type_=PluginTextEmbeddingNumTokensResponse,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -334,7 +334,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/rerank/invoke",
|
path=f"plugin/{tenant_id}/dispatch/rerank/invoke",
|
||||||
type=RerankResult,
|
type_=RerankResult,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -378,7 +378,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/tts/invoke",
|
path=f"plugin/{tenant_id}/dispatch/tts/invoke",
|
||||||
type=PluginStringResultResponse,
|
type_=PluginStringResultResponse,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -422,7 +422,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/tts/model/voices",
|
path=f"plugin/{tenant_id}/dispatch/tts/model/voices",
|
||||||
type=PluginVoicesResponse,
|
type_=PluginVoicesResponse,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -466,7 +466,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/speech2text/invoke",
|
path=f"plugin/{tenant_id}/dispatch/speech2text/invoke",
|
||||||
type=PluginStringResultResponse,
|
type_=PluginStringResultResponse,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
@ -506,7 +506,7 @@ class PluginModelClient(BasePluginClient):
|
|||||||
response = self._request_with_plugin_daemon_response_stream(
|
response = self._request_with_plugin_daemon_response_stream(
|
||||||
method="POST",
|
method="POST",
|
||||||
path=f"plugin/{tenant_id}/dispatch/moderation/invoke",
|
path=f"plugin/{tenant_id}/dispatch/moderation/invoke",
|
||||||
type=PluginBasicBooleanResponse,
|
type_=PluginBasicBooleanResponse,
|
||||||
data=jsonable_encoder(
|
data=jsonable_encoder(
|
||||||
{
|
{
|
||||||
"user_id": user_id,
|
"user_id": user_id,
|
||||||
|
|||||||
@ -610,7 +610,7 @@ class ProviderManager:
|
|||||||
|
|
||||||
provider_quota_to_provider_record_dict = {}
|
provider_quota_to_provider_record_dict = {}
|
||||||
for provider_record in provider_records:
|
for provider_record in provider_records:
|
||||||
if provider_record.provider_type != ProviderType.SYSTEM.value:
|
if provider_record.provider_type != ProviderType.SYSTEM:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
provider_quota_to_provider_record_dict[ProviderQuotaType.value_of(provider_record.quota_type)] = (
|
provider_quota_to_provider_record_dict[ProviderQuotaType.value_of(provider_record.quota_type)] = (
|
||||||
@ -627,8 +627,8 @@ class ProviderManager:
|
|||||||
tenant_id=tenant_id,
|
tenant_id=tenant_id,
|
||||||
# TODO: Use provider name with prefix after the data migration.
|
# TODO: Use provider name with prefix after the data migration.
|
||||||
provider_name=ModelProviderID(provider_name).provider_name,
|
provider_name=ModelProviderID(provider_name).provider_name,
|
||||||
provider_type=ProviderType.SYSTEM.value,
|
provider_type=ProviderType.SYSTEM,
|
||||||
quota_type=ProviderQuotaType.TRIAL.value,
|
quota_type=ProviderQuotaType.TRIAL,
|
||||||
quota_limit=quota.quota_limit, # type: ignore
|
quota_limit=quota.quota_limit, # type: ignore
|
||||||
quota_used=0,
|
quota_used=0,
|
||||||
is_valid=True,
|
is_valid=True,
|
||||||
@ -641,8 +641,8 @@ class ProviderManager:
|
|||||||
stmt = select(Provider).where(
|
stmt = select(Provider).where(
|
||||||
Provider.tenant_id == tenant_id,
|
Provider.tenant_id == tenant_id,
|
||||||
Provider.provider_name == ModelProviderID(provider_name).provider_name,
|
Provider.provider_name == ModelProviderID(provider_name).provider_name,
|
||||||
Provider.provider_type == ProviderType.SYSTEM.value,
|
Provider.provider_type == ProviderType.SYSTEM,
|
||||||
Provider.quota_type == ProviderQuotaType.TRIAL.value,
|
Provider.quota_type == ProviderQuotaType.TRIAL,
|
||||||
)
|
)
|
||||||
existed_provider_record = db.session.scalar(stmt)
|
existed_provider_record = db.session.scalar(stmt)
|
||||||
if not existed_provider_record:
|
if not existed_provider_record:
|
||||||
@ -702,7 +702,7 @@ class ProviderManager:
|
|||||||
"""Get custom provider configuration."""
|
"""Get custom provider configuration."""
|
||||||
# Find custom provider record (non-system)
|
# Find custom provider record (non-system)
|
||||||
custom_provider_record = next(
|
custom_provider_record = next(
|
||||||
(record for record in provider_records if record.provider_type != ProviderType.SYSTEM.value), None
|
(record for record in provider_records if record.provider_type != ProviderType.SYSTEM), None
|
||||||
)
|
)
|
||||||
|
|
||||||
if not custom_provider_record:
|
if not custom_provider_record:
|
||||||
@ -905,7 +905,7 @@ class ProviderManager:
|
|||||||
# Convert provider_records to dict
|
# Convert provider_records to dict
|
||||||
quota_type_to_provider_records_dict: dict[ProviderQuotaType, Provider] = {}
|
quota_type_to_provider_records_dict: dict[ProviderQuotaType, Provider] = {}
|
||||||
for provider_record in provider_records:
|
for provider_record in provider_records:
|
||||||
if provider_record.provider_type != ProviderType.SYSTEM.value:
|
if provider_record.provider_type != ProviderType.SYSTEM:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
quota_type_to_provider_records_dict[ProviderQuotaType.value_of(provider_record.quota_type)] = (
|
quota_type_to_provider_records_dict[ProviderQuotaType.value_of(provider_record.quota_type)] = (
|
||||||
@ -1046,7 +1046,7 @@ class ProviderManager:
|
|||||||
"""
|
"""
|
||||||
secret_input_form_variables = []
|
secret_input_form_variables = []
|
||||||
for credential_form_schema in credential_form_schemas:
|
for credential_form_schema in credential_form_schemas:
|
||||||
if credential_form_schema.type.value == FormType.SECRET_INPUT.value:
|
if credential_form_schema.type == FormType.SECRET_INPUT:
|
||||||
secret_input_form_variables.append(credential_form_schema.variable)
|
secret_input_form_variables.append(credential_form_schema.variable)
|
||||||
|
|
||||||
return secret_input_form_variables
|
return secret_input_form_variables
|
||||||
|
|||||||
@ -46,7 +46,7 @@ class DataPostProcessor:
|
|||||||
reranking_model: dict | None = None,
|
reranking_model: dict | None = None,
|
||||||
weights: dict | None = None,
|
weights: dict | None = None,
|
||||||
) -> BaseRerankRunner | None:
|
) -> BaseRerankRunner | None:
|
||||||
if reranking_mode == RerankMode.WEIGHTED_SCORE.value and weights:
|
if reranking_mode == RerankMode.WEIGHTED_SCORE and weights:
|
||||||
runner = RerankRunnerFactory.create_rerank_runner(
|
runner = RerankRunnerFactory.create_rerank_runner(
|
||||||
runner_type=reranking_mode,
|
runner_type=reranking_mode,
|
||||||
tenant_id=tenant_id,
|
tenant_id=tenant_id,
|
||||||
@ -62,7 +62,7 @@ class DataPostProcessor:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
return runner
|
return runner
|
||||||
elif reranking_mode == RerankMode.RERANKING_MODEL.value:
|
elif reranking_mode == RerankMode.RERANKING_MODEL:
|
||||||
rerank_model_instance = self._get_rerank_model_instance(tenant_id, reranking_model)
|
rerank_model_instance = self._get_rerank_model_instance(tenant_id, reranking_model)
|
||||||
if rerank_model_instance is None:
|
if rerank_model_instance is None:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@ -21,7 +21,7 @@ from models.dataset import Document as DatasetDocument
|
|||||||
from services.external_knowledge_service import ExternalDatasetService
|
from services.external_knowledge_service import ExternalDatasetService
|
||||||
|
|
||||||
default_retrieval_model = {
|
default_retrieval_model = {
|
||||||
"search_method": RetrievalMethod.SEMANTIC_SEARCH.value,
|
"search_method": RetrievalMethod.SEMANTIC_SEARCH,
|
||||||
"reranking_enable": False,
|
"reranking_enable": False,
|
||||||
"reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""},
|
"reranking_model": {"reranking_provider_name": "", "reranking_model_name": ""},
|
||||||
"top_k": 4,
|
"top_k": 4,
|
||||||
@ -34,7 +34,7 @@ class RetrievalService:
|
|||||||
@classmethod
|
@classmethod
|
||||||
def retrieve(
|
def retrieve(
|
||||||
cls,
|
cls,
|
||||||
retrieval_method: str,
|
retrieval_method: RetrievalMethod,
|
||||||
dataset_id: str,
|
dataset_id: str,
|
||||||
query: str,
|
query: str,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
@ -56,7 +56,7 @@ class RetrievalService:
|
|||||||
# Optimize multithreading with thread pools
|
# Optimize multithreading with thread pools
|
||||||
with ThreadPoolExecutor(max_workers=dify_config.RETRIEVAL_SERVICE_EXECUTORS) as executor: # type: ignore
|
with ThreadPoolExecutor(max_workers=dify_config.RETRIEVAL_SERVICE_EXECUTORS) as executor: # type: ignore
|
||||||
futures = []
|
futures = []
|
||||||
if retrieval_method == "keyword_search":
|
if retrieval_method == RetrievalMethod.KEYWORD_SEARCH:
|
||||||
futures.append(
|
futures.append(
|
||||||
executor.submit(
|
executor.submit(
|
||||||
cls.keyword_search,
|
cls.keyword_search,
|
||||||
@ -107,7 +107,7 @@ class RetrievalService:
|
|||||||
raise ValueError(";\n".join(exceptions))
|
raise ValueError(";\n".join(exceptions))
|
||||||
|
|
||||||
# Deduplicate documents for hybrid search to avoid duplicate chunks
|
# Deduplicate documents for hybrid search to avoid duplicate chunks
|
||||||
if retrieval_method == RetrievalMethod.HYBRID_SEARCH.value:
|
if retrieval_method == RetrievalMethod.HYBRID_SEARCH:
|
||||||
all_documents = cls._deduplicate_documents(all_documents)
|
all_documents = cls._deduplicate_documents(all_documents)
|
||||||
data_post_processor = DataPostProcessor(
|
data_post_processor = DataPostProcessor(
|
||||||
str(dataset.tenant_id), reranking_mode, reranking_model, weights, False
|
str(dataset.tenant_id), reranking_mode, reranking_model, weights, False
|
||||||
@ -134,7 +134,7 @@ class RetrievalService:
|
|||||||
if not dataset:
|
if not dataset:
|
||||||
return []
|
return []
|
||||||
metadata_condition = (
|
metadata_condition = (
|
||||||
MetadataCondition(**metadata_filtering_conditions) if metadata_filtering_conditions else None
|
MetadataCondition.model_validate(metadata_filtering_conditions) if metadata_filtering_conditions else None
|
||||||
)
|
)
|
||||||
all_documents = ExternalDatasetService.fetch_external_knowledge_retrieval(
|
all_documents = ExternalDatasetService.fetch_external_knowledge_retrieval(
|
||||||
dataset.tenant_id,
|
dataset.tenant_id,
|
||||||
@ -220,7 +220,7 @@ class RetrievalService:
|
|||||||
score_threshold: float | None,
|
score_threshold: float | None,
|
||||||
reranking_model: dict | None,
|
reranking_model: dict | None,
|
||||||
all_documents: list,
|
all_documents: list,
|
||||||
retrieval_method: str,
|
retrieval_method: RetrievalMethod,
|
||||||
exceptions: list,
|
exceptions: list,
|
||||||
document_ids_filter: list[str] | None = None,
|
document_ids_filter: list[str] | None = None,
|
||||||
):
|
):
|
||||||
@ -245,10 +245,10 @@ class RetrievalService:
|
|||||||
reranking_model
|
reranking_model
|
||||||
and reranking_model.get("reranking_model_name")
|
and reranking_model.get("reranking_model_name")
|
||||||
and reranking_model.get("reranking_provider_name")
|
and reranking_model.get("reranking_provider_name")
|
||||||
and retrieval_method == RetrievalMethod.SEMANTIC_SEARCH.value
|
and retrieval_method == RetrievalMethod.SEMANTIC_SEARCH
|
||||||
):
|
):
|
||||||
data_post_processor = DataPostProcessor(
|
data_post_processor = DataPostProcessor(
|
||||||
str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL.value), reranking_model, None, False
|
str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL), reranking_model, None, False
|
||||||
)
|
)
|
||||||
all_documents.extend(
|
all_documents.extend(
|
||||||
data_post_processor.invoke(
|
data_post_processor.invoke(
|
||||||
@ -293,10 +293,10 @@ class RetrievalService:
|
|||||||
reranking_model
|
reranking_model
|
||||||
and reranking_model.get("reranking_model_name")
|
and reranking_model.get("reranking_model_name")
|
||||||
and reranking_model.get("reranking_provider_name")
|
and reranking_model.get("reranking_provider_name")
|
||||||
and retrieval_method == RetrievalMethod.FULL_TEXT_SEARCH.value
|
and retrieval_method == RetrievalMethod.FULL_TEXT_SEARCH
|
||||||
):
|
):
|
||||||
data_post_processor = DataPostProcessor(
|
data_post_processor = DataPostProcessor(
|
||||||
str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL.value), reranking_model, None, False
|
str(dataset.tenant_id), str(RerankMode.RERANKING_MODEL), reranking_model, None, False
|
||||||
)
|
)
|
||||||
all_documents.extend(
|
all_documents.extend(
|
||||||
data_post_processor.invoke(
|
data_post_processor.invoke(
|
||||||
|
|||||||
@ -0,0 +1,388 @@
|
|||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from contextlib import contextmanager
|
||||||
|
from typing import Any, Literal, cast
|
||||||
|
|
||||||
|
import mysql.connector
|
||||||
|
from mysql.connector import Error as MySQLError
|
||||||
|
from pydantic import BaseModel, model_validator
|
||||||
|
|
||||||
|
from configs import dify_config
|
||||||
|
from core.rag.datasource.vdb.vector_base import BaseVector
|
||||||
|
from core.rag.datasource.vdb.vector_factory import AbstractVectorFactory
|
||||||
|
from core.rag.datasource.vdb.vector_type import VectorType
|
||||||
|
from core.rag.embedding.embedding_base import Embeddings
|
||||||
|
from core.rag.models.document import Document
|
||||||
|
from extensions.ext_redis import redis_client
|
||||||
|
from models.dataset import Dataset
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class AlibabaCloudMySQLVectorConfig(BaseModel):
|
||||||
|
host: str
|
||||||
|
port: int
|
||||||
|
user: str
|
||||||
|
password: str
|
||||||
|
database: str
|
||||||
|
max_connection: int
|
||||||
|
charset: str = "utf8mb4"
|
||||||
|
distance_function: Literal["cosine", "euclidean"] = "cosine"
|
||||||
|
hnsw_m: int = 6
|
||||||
|
|
||||||
|
@model_validator(mode="before")
|
||||||
|
@classmethod
|
||||||
|
def validate_config(cls, values: dict):
|
||||||
|
if not values.get("host"):
|
||||||
|
raise ValueError("config ALIBABACLOUD_MYSQL_HOST is required")
|
||||||
|
if not values.get("port"):
|
||||||
|
raise ValueError("config ALIBABACLOUD_MYSQL_PORT is required")
|
||||||
|
if not values.get("user"):
|
||||||
|
raise ValueError("config ALIBABACLOUD_MYSQL_USER is required")
|
||||||
|
if values.get("password") is None:
|
||||||
|
raise ValueError("config ALIBABACLOUD_MYSQL_PASSWORD is required")
|
||||||
|
if not values.get("database"):
|
||||||
|
raise ValueError("config ALIBABACLOUD_MYSQL_DATABASE is required")
|
||||||
|
if not values.get("max_connection"):
|
||||||
|
raise ValueError("config ALIBABACLOUD_MYSQL_MAX_CONNECTION is required")
|
||||||
|
return values
|
||||||
|
|
||||||
|
|
||||||
|
SQL_CREATE_TABLE = """
|
||||||
|
CREATE TABLE IF NOT EXISTS {table_name} (
|
||||||
|
id VARCHAR(36) PRIMARY KEY,
|
||||||
|
text LONGTEXT NOT NULL,
|
||||||
|
meta JSON NOT NULL,
|
||||||
|
embedding VECTOR({dimension}) NOT NULL,
|
||||||
|
VECTOR INDEX (embedding) M={hnsw_m} DISTANCE={distance_function}
|
||||||
|
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci;
|
||||||
|
"""
|
||||||
|
|
||||||
|
SQL_CREATE_META_INDEX = """
|
||||||
|
CREATE INDEX idx_{index_hash}_meta ON {table_name}
|
||||||
|
((CAST(JSON_UNQUOTE(JSON_EXTRACT(meta, '$.document_id')) AS CHAR(36))));
|
||||||
|
"""
|
||||||
|
|
||||||
|
SQL_CREATE_FULLTEXT_INDEX = """
|
||||||
|
CREATE FULLTEXT INDEX idx_{index_hash}_text ON {table_name} (text) WITH PARSER ngram;
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class AlibabaCloudMySQLVector(BaseVector):
|
||||||
|
def __init__(self, collection_name: str, config: AlibabaCloudMySQLVectorConfig):
|
||||||
|
super().__init__(collection_name)
|
||||||
|
self.pool = self._create_connection_pool(config)
|
||||||
|
self.table_name = collection_name.lower()
|
||||||
|
self.index_hash = hashlib.md5(self.table_name.encode()).hexdigest()[:8]
|
||||||
|
self.distance_function = config.distance_function.lower()
|
||||||
|
self.hnsw_m = config.hnsw_m
|
||||||
|
self._check_vector_support()
|
||||||
|
|
||||||
|
def get_type(self) -> str:
|
||||||
|
return VectorType.ALIBABACLOUD_MYSQL
|
||||||
|
|
||||||
|
def _create_connection_pool(self, config: AlibabaCloudMySQLVectorConfig):
|
||||||
|
# Create connection pool using mysql-connector-python pooling
|
||||||
|
pool_config: dict[str, Any] = {
|
||||||
|
"host": config.host,
|
||||||
|
"port": config.port,
|
||||||
|
"user": config.user,
|
||||||
|
"password": config.password,
|
||||||
|
"database": config.database,
|
||||||
|
"charset": config.charset,
|
||||||
|
"autocommit": True,
|
||||||
|
"pool_name": f"pool_{self.collection_name}",
|
||||||
|
"pool_size": config.max_connection,
|
||||||
|
"pool_reset_session": True,
|
||||||
|
}
|
||||||
|
return mysql.connector.pooling.MySQLConnectionPool(**pool_config)
|
||||||
|
|
||||||
|
def _check_vector_support(self):
|
||||||
|
"""Check if the MySQL server supports vector operations."""
|
||||||
|
try:
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
# Check MySQL version and vector support
|
||||||
|
cur.execute("SELECT VERSION()")
|
||||||
|
version = cur.fetchone()["VERSION()"]
|
||||||
|
logger.debug("Connected to MySQL version: %s", version)
|
||||||
|
# Try to execute a simple vector function to verify support
|
||||||
|
cur.execute("SELECT VEC_FromText('[1,2,3]') IS NOT NULL as vector_support")
|
||||||
|
result = cur.fetchone()
|
||||||
|
if not result or not result.get("vector_support"):
|
||||||
|
raise ValueError(
|
||||||
|
"RDS MySQL Vector functions are not available."
|
||||||
|
" Please ensure you're using RDS MySQL 8.0.36+ with Vector support."
|
||||||
|
)
|
||||||
|
|
||||||
|
except MySQLError as e:
|
||||||
|
if "FUNCTION" in str(e) and "VEC_FromText" in str(e):
|
||||||
|
raise ValueError(
|
||||||
|
"RDS MySQL Vector functions are not available."
|
||||||
|
" Please ensure you're using RDS MySQL 8.0.36+ with Vector support."
|
||||||
|
) from e
|
||||||
|
raise e
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _get_cursor(self):
|
||||||
|
conn = self.pool.get_connection()
|
||||||
|
cur = conn.cursor(dictionary=True)
|
||||||
|
try:
|
||||||
|
yield cur
|
||||||
|
finally:
|
||||||
|
cur.close()
|
||||||
|
conn.close()
|
||||||
|
|
||||||
|
def create(self, texts: list[Document], embeddings: list[list[float]], **kwargs):
|
||||||
|
dimension = len(embeddings[0])
|
||||||
|
self._create_collection(dimension)
|
||||||
|
return self.add_texts(texts, embeddings)
|
||||||
|
|
||||||
|
def add_texts(self, documents: list[Document], embeddings: list[list[float]], **kwargs):
|
||||||
|
values = []
|
||||||
|
pks = []
|
||||||
|
for i, doc in enumerate(documents):
|
||||||
|
if doc.metadata is not None:
|
||||||
|
doc_id = doc.metadata.get("doc_id", str(uuid.uuid4()))
|
||||||
|
pks.append(doc_id)
|
||||||
|
# Convert embedding list to Aliyun MySQL vector format
|
||||||
|
vector_str = "[" + ",".join(map(str, embeddings[i])) + "]"
|
||||||
|
values.append(
|
||||||
|
(
|
||||||
|
doc_id,
|
||||||
|
doc.page_content,
|
||||||
|
json.dumps(doc.metadata),
|
||||||
|
vector_str,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
insert_sql = (
|
||||||
|
f"INSERT INTO {self.table_name} (id, text, meta, embedding) VALUES (%s, %s, %s, VEC_FromText(%s))"
|
||||||
|
)
|
||||||
|
cur.executemany(insert_sql, values)
|
||||||
|
return pks
|
||||||
|
|
||||||
|
def text_exists(self, id: str) -> bool:
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
cur.execute(f"SELECT id FROM {self.table_name} WHERE id = %s", (id,))
|
||||||
|
return cur.fetchone() is not None
|
||||||
|
|
||||||
|
def get_by_ids(self, ids: list[str]) -> list[Document]:
|
||||||
|
if not ids:
|
||||||
|
return []
|
||||||
|
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
placeholders = ",".join(["%s"] * len(ids))
|
||||||
|
cur.execute(f"SELECT meta, text FROM {self.table_name} WHERE id IN ({placeholders})", ids)
|
||||||
|
docs = []
|
||||||
|
for record in cur:
|
||||||
|
metadata = record["meta"]
|
||||||
|
if isinstance(metadata, str):
|
||||||
|
metadata = json.loads(metadata)
|
||||||
|
docs.append(Document(page_content=record["text"], metadata=metadata))
|
||||||
|
return docs
|
||||||
|
|
||||||
|
def delete_by_ids(self, ids: list[str]):
|
||||||
|
# Avoiding crashes caused by performing delete operations on empty lists
|
||||||
|
if not ids:
|
||||||
|
return
|
||||||
|
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
try:
|
||||||
|
placeholders = ",".join(["%s"] * len(ids))
|
||||||
|
cur.execute(f"DELETE FROM {self.table_name} WHERE id IN ({placeholders})", ids)
|
||||||
|
except MySQLError as e:
|
||||||
|
if e.errno == 1146: # Table doesn't exist
|
||||||
|
logger.warning("Table %s not found, skipping delete operation.", self.table_name)
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
raise e
|
||||||
|
|
||||||
|
def delete_by_metadata_field(self, key: str, value: str):
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
cur.execute(
|
||||||
|
f"DELETE FROM {self.table_name} WHERE JSON_UNQUOTE(JSON_EXTRACT(meta, %s)) = %s", (f"$.{key}", value)
|
||||||
|
)
|
||||||
|
|
||||||
|
def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:
|
||||||
|
"""
|
||||||
|
Search the nearest neighbors to a vector using RDS MySQL vector distance functions.
|
||||||
|
|
||||||
|
:param query_vector: The input vector to search for similar items.
|
||||||
|
:return: List of Documents that are nearest to the query vector.
|
||||||
|
"""
|
||||||
|
top_k = kwargs.get("top_k", 4)
|
||||||
|
if not isinstance(top_k, int) or top_k <= 0:
|
||||||
|
raise ValueError("top_k must be a positive integer")
|
||||||
|
|
||||||
|
document_ids_filter = kwargs.get("document_ids_filter")
|
||||||
|
where_clause = ""
|
||||||
|
params = []
|
||||||
|
|
||||||
|
if document_ids_filter:
|
||||||
|
placeholders = ",".join(["%s"] * len(document_ids_filter))
|
||||||
|
where_clause = f" WHERE JSON_UNQUOTE(JSON_EXTRACT(meta, '$.document_id')) IN ({placeholders}) "
|
||||||
|
params.extend(document_ids_filter)
|
||||||
|
|
||||||
|
# Convert query vector to RDS MySQL vector format
|
||||||
|
query_vector_str = "[" + ",".join(map(str, query_vector)) + "]"
|
||||||
|
|
||||||
|
# Use RSD MySQL's native vector distance functions
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
# Choose distance function based on configuration
|
||||||
|
distance_func = "VEC_DISTANCE_COSINE" if self.distance_function == "cosine" else "VEC_DISTANCE_EUCLIDEAN"
|
||||||
|
|
||||||
|
# Note: RDS MySQL optimizer will use vector index when ORDER BY + LIMIT are present
|
||||||
|
# Use column alias in ORDER BY to avoid calculating distance twice
|
||||||
|
sql = f"""
|
||||||
|
SELECT meta, text,
|
||||||
|
{distance_func}(embedding, VEC_FromText(%s)) AS distance
|
||||||
|
FROM {self.table_name}
|
||||||
|
{where_clause}
|
||||||
|
ORDER BY distance
|
||||||
|
LIMIT %s
|
||||||
|
"""
|
||||||
|
query_params = [query_vector_str] + params + [top_k]
|
||||||
|
|
||||||
|
cur.execute(sql, query_params)
|
||||||
|
|
||||||
|
docs = []
|
||||||
|
score_threshold = float(kwargs.get("score_threshold") or 0.0)
|
||||||
|
|
||||||
|
for record in cur:
|
||||||
|
try:
|
||||||
|
distance = float(record["distance"])
|
||||||
|
# Convert distance to similarity score
|
||||||
|
if self.distance_function == "cosine":
|
||||||
|
# For cosine distance: similarity = 1 - distance
|
||||||
|
similarity = 1.0 - distance
|
||||||
|
else:
|
||||||
|
# For euclidean distance: use inverse relationship
|
||||||
|
# similarity = 1 / (1 + distance)
|
||||||
|
similarity = 1.0 / (1.0 + distance)
|
||||||
|
|
||||||
|
metadata = record["meta"]
|
||||||
|
if isinstance(metadata, str):
|
||||||
|
metadata = json.loads(metadata)
|
||||||
|
metadata["score"] = similarity
|
||||||
|
metadata["distance"] = distance
|
||||||
|
|
||||||
|
if similarity >= score_threshold:
|
||||||
|
docs.append(Document(page_content=record["text"], metadata=metadata))
|
||||||
|
except (ValueError, json.JSONDecodeError) as e:
|
||||||
|
logger.warning("Error processing search result: %s", e)
|
||||||
|
continue
|
||||||
|
|
||||||
|
return docs
|
||||||
|
|
||||||
|
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
||||||
|
top_k = kwargs.get("top_k", 5)
|
||||||
|
if not isinstance(top_k, int) or top_k <= 0:
|
||||||
|
raise ValueError("top_k must be a positive integer")
|
||||||
|
|
||||||
|
document_ids_filter = kwargs.get("document_ids_filter")
|
||||||
|
where_clause = ""
|
||||||
|
params = []
|
||||||
|
|
||||||
|
if document_ids_filter:
|
||||||
|
placeholders = ",".join(["%s"] * len(document_ids_filter))
|
||||||
|
where_clause = f" AND JSON_UNQUOTE(JSON_EXTRACT(meta, '$.document_id')) IN ({placeholders}) "
|
||||||
|
params.extend(document_ids_filter)
|
||||||
|
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
# Build query parameters: query (twice for MATCH clauses), document_ids_filter (if any), top_k
|
||||||
|
query_params = [query, query] + params + [top_k]
|
||||||
|
cur.execute(
|
||||||
|
f"""SELECT meta, text,
|
||||||
|
MATCH(text) AGAINST(%s IN NATURAL LANGUAGE MODE) AS score
|
||||||
|
FROM {self.table_name}
|
||||||
|
WHERE MATCH(text) AGAINST(%s IN NATURAL LANGUAGE MODE)
|
||||||
|
{where_clause}
|
||||||
|
ORDER BY score DESC
|
||||||
|
LIMIT %s""",
|
||||||
|
query_params,
|
||||||
|
)
|
||||||
|
docs = []
|
||||||
|
for record in cur:
|
||||||
|
metadata = record["meta"]
|
||||||
|
if isinstance(metadata, str):
|
||||||
|
metadata = json.loads(metadata)
|
||||||
|
metadata["score"] = float(record["score"])
|
||||||
|
docs.append(Document(page_content=record["text"], metadata=metadata))
|
||||||
|
return docs
|
||||||
|
|
||||||
|
def delete(self):
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
cur.execute(f"DROP TABLE IF EXISTS {self.table_name}")
|
||||||
|
|
||||||
|
def _create_collection(self, dimension: int):
|
||||||
|
collection_exist_cache_key = f"vector_indexing_{self._collection_name}"
|
||||||
|
lock_name = f"{collection_exist_cache_key}_lock"
|
||||||
|
with redis_client.lock(lock_name, timeout=20):
|
||||||
|
if redis_client.get(collection_exist_cache_key):
|
||||||
|
return
|
||||||
|
|
||||||
|
with self._get_cursor() as cur:
|
||||||
|
# Create table with vector column and vector index
|
||||||
|
cur.execute(
|
||||||
|
SQL_CREATE_TABLE.format(
|
||||||
|
table_name=self.table_name,
|
||||||
|
dimension=dimension,
|
||||||
|
distance_function=self.distance_function,
|
||||||
|
hnsw_m=self.hnsw_m,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# Create metadata index (check if exists first)
|
||||||
|
try:
|
||||||
|
cur.execute(SQL_CREATE_META_INDEX.format(table_name=self.table_name, index_hash=self.index_hash))
|
||||||
|
except MySQLError as e:
|
||||||
|
if e.errno != 1061: # Duplicate key name
|
||||||
|
logger.warning("Could not create meta index: %s", e)
|
||||||
|
|
||||||
|
# Create full-text index for text search
|
||||||
|
try:
|
||||||
|
cur.execute(
|
||||||
|
SQL_CREATE_FULLTEXT_INDEX.format(table_name=self.table_name, index_hash=self.index_hash)
|
||||||
|
)
|
||||||
|
except MySQLError as e:
|
||||||
|
if e.errno != 1061: # Duplicate key name
|
||||||
|
logger.warning("Could not create fulltext index: %s", e)
|
||||||
|
|
||||||
|
redis_client.set(collection_exist_cache_key, 1, ex=3600)
|
||||||
|
|
||||||
|
|
||||||
|
class AlibabaCloudMySQLVectorFactory(AbstractVectorFactory):
|
||||||
|
def _validate_distance_function(self, distance_function: str) -> Literal["cosine", "euclidean"]:
|
||||||
|
"""Validate and return the distance function as a proper Literal type."""
|
||||||
|
if distance_function not in ["cosine", "euclidean"]:
|
||||||
|
raise ValueError(f"Invalid distance function: {distance_function}. Must be 'cosine' or 'euclidean'")
|
||||||
|
return cast(Literal["cosine", "euclidean"], distance_function)
|
||||||
|
|
||||||
|
def init_vector(self, dataset: Dataset, attributes: list, embeddings: Embeddings) -> AlibabaCloudMySQLVector:
|
||||||
|
if dataset.index_struct_dict:
|
||||||
|
class_prefix: str = dataset.index_struct_dict["vector_store"]["class_prefix"]
|
||||||
|
collection_name = class_prefix
|
||||||
|
else:
|
||||||
|
dataset_id = dataset.id
|
||||||
|
collection_name = Dataset.gen_collection_name_by_id(dataset_id)
|
||||||
|
dataset.index_struct = json.dumps(
|
||||||
|
self.gen_index_struct_dict(VectorType.ALIBABACLOUD_MYSQL, collection_name)
|
||||||
|
)
|
||||||
|
return AlibabaCloudMySQLVector(
|
||||||
|
collection_name=collection_name,
|
||||||
|
config=AlibabaCloudMySQLVectorConfig(
|
||||||
|
host=dify_config.ALIBABACLOUD_MYSQL_HOST or "localhost",
|
||||||
|
port=dify_config.ALIBABACLOUD_MYSQL_PORT,
|
||||||
|
user=dify_config.ALIBABACLOUD_MYSQL_USER or "root",
|
||||||
|
password=dify_config.ALIBABACLOUD_MYSQL_PASSWORD or "",
|
||||||
|
database=dify_config.ALIBABACLOUD_MYSQL_DATABASE or "dify",
|
||||||
|
max_connection=dify_config.ALIBABACLOUD_MYSQL_MAX_CONNECTION,
|
||||||
|
charset=dify_config.ALIBABACLOUD_MYSQL_CHARSET or "utf8mb4",
|
||||||
|
distance_function=self._validate_distance_function(
|
||||||
|
dify_config.ALIBABACLOUD_MYSQL_DISTANCE_FUNCTION or "cosine"
|
||||||
|
),
|
||||||
|
hnsw_m=dify_config.ALIBABACLOUD_MYSQL_HNSW_M or 6,
|
||||||
|
),
|
||||||
|
)
|
||||||
@ -488,9 +488,9 @@ class ClickzettaVector(BaseVector):
|
|||||||
create_table_sql = f"""
|
create_table_sql = f"""
|
||||||
CREATE TABLE IF NOT EXISTS {self._config.schema_name}.{self._table_name} (
|
CREATE TABLE IF NOT EXISTS {self._config.schema_name}.{self._table_name} (
|
||||||
id STRING NOT NULL COMMENT 'Unique document identifier',
|
id STRING NOT NULL COMMENT 'Unique document identifier',
|
||||||
{Field.CONTENT_KEY.value} STRING NOT NULL COMMENT 'Document text content for search and retrieval',
|
{Field.CONTENT_KEY} STRING NOT NULL COMMENT 'Document text content for search and retrieval',
|
||||||
{Field.METADATA_KEY.value} JSON COMMENT 'Document metadata including source, type, and other attributes',
|
{Field.METADATA_KEY} JSON COMMENT 'Document metadata including source, type, and other attributes',
|
||||||
{Field.VECTOR.value} VECTOR(FLOAT, {dimension}) NOT NULL COMMENT
|
{Field.VECTOR} VECTOR(FLOAT, {dimension}) NOT NULL COMMENT
|
||||||
'High-dimensional embedding vector for semantic similarity search',
|
'High-dimensional embedding vector for semantic similarity search',
|
||||||
PRIMARY KEY (id)
|
PRIMARY KEY (id)
|
||||||
) COMMENT 'Dify RAG knowledge base vector storage table for document embeddings and content'
|
) COMMENT 'Dify RAG knowledge base vector storage table for document embeddings and content'
|
||||||
@ -519,15 +519,15 @@ class ClickzettaVector(BaseVector):
|
|||||||
existing_indexes = cursor.fetchall()
|
existing_indexes = cursor.fetchall()
|
||||||
for idx in existing_indexes:
|
for idx in existing_indexes:
|
||||||
# Check if vector index already exists on the embedding column
|
# Check if vector index already exists on the embedding column
|
||||||
if Field.VECTOR.value in str(idx).lower():
|
if Field.VECTOR in str(idx).lower():
|
||||||
logger.info("Vector index already exists on column %s", Field.VECTOR.value)
|
logger.info("Vector index already exists on column %s", Field.VECTOR)
|
||||||
return
|
return
|
||||||
except (RuntimeError, ValueError) as e:
|
except (RuntimeError, ValueError) as e:
|
||||||
logger.warning("Failed to check existing indexes: %s", e)
|
logger.warning("Failed to check existing indexes: %s", e)
|
||||||
|
|
||||||
index_sql = f"""
|
index_sql = f"""
|
||||||
CREATE VECTOR INDEX IF NOT EXISTS {index_name}
|
CREATE VECTOR INDEX IF NOT EXISTS {index_name}
|
||||||
ON TABLE {self._config.schema_name}.{self._table_name}({Field.VECTOR.value})
|
ON TABLE {self._config.schema_name}.{self._table_name}({Field.VECTOR})
|
||||||
PROPERTIES (
|
PROPERTIES (
|
||||||
"distance.function" = "{self._config.vector_distance_function}",
|
"distance.function" = "{self._config.vector_distance_function}",
|
||||||
"scalar.type" = "f32",
|
"scalar.type" = "f32",
|
||||||
@ -560,17 +560,17 @@ class ClickzettaVector(BaseVector):
|
|||||||
# More precise check: look for inverted index specifically on the content column
|
# More precise check: look for inverted index specifically on the content column
|
||||||
if (
|
if (
|
||||||
"inverted" in idx_str
|
"inverted" in idx_str
|
||||||
and Field.CONTENT_KEY.value.lower() in idx_str
|
and Field.CONTENT_KEY.lower() in idx_str
|
||||||
and (index_name.lower() in idx_str or f"idx_{self._table_name}_text" in idx_str)
|
and (index_name.lower() in idx_str or f"idx_{self._table_name}_text" in idx_str)
|
||||||
):
|
):
|
||||||
logger.info("Inverted index already exists on column %s: %s", Field.CONTENT_KEY.value, idx)
|
logger.info("Inverted index already exists on column %s: %s", Field.CONTENT_KEY, idx)
|
||||||
return
|
return
|
||||||
except (RuntimeError, ValueError) as e:
|
except (RuntimeError, ValueError) as e:
|
||||||
logger.warning("Failed to check existing indexes: %s", e)
|
logger.warning("Failed to check existing indexes: %s", e)
|
||||||
|
|
||||||
index_sql = f"""
|
index_sql = f"""
|
||||||
CREATE INVERTED INDEX IF NOT EXISTS {index_name}
|
CREATE INVERTED INDEX IF NOT EXISTS {index_name}
|
||||||
ON TABLE {self._config.schema_name}.{self._table_name} ({Field.CONTENT_KEY.value})
|
ON TABLE {self._config.schema_name}.{self._table_name} ({Field.CONTENT_KEY})
|
||||||
PROPERTIES (
|
PROPERTIES (
|
||||||
"analyzer" = "{self._config.analyzer_type}",
|
"analyzer" = "{self._config.analyzer_type}",
|
||||||
"mode" = "{self._config.analyzer_mode}"
|
"mode" = "{self._config.analyzer_mode}"
|
||||||
@ -588,13 +588,13 @@ class ClickzettaVector(BaseVector):
|
|||||||
or "with the same type" in error_msg
|
or "with the same type" in error_msg
|
||||||
or "cannot create inverted index" in error_msg
|
or "cannot create inverted index" in error_msg
|
||||||
) and "already has index" in error_msg:
|
) and "already has index" in error_msg:
|
||||||
logger.info("Inverted index already exists on column %s", Field.CONTENT_KEY.value)
|
logger.info("Inverted index already exists on column %s", Field.CONTENT_KEY)
|
||||||
# Try to get the existing index name for logging
|
# Try to get the existing index name for logging
|
||||||
try:
|
try:
|
||||||
cursor.execute(f"SHOW INDEX FROM {self._config.schema_name}.{self._table_name}")
|
cursor.execute(f"SHOW INDEX FROM {self._config.schema_name}.{self._table_name}")
|
||||||
existing_indexes = cursor.fetchall()
|
existing_indexes = cursor.fetchall()
|
||||||
for idx in existing_indexes:
|
for idx in existing_indexes:
|
||||||
if "inverted" in str(idx).lower() and Field.CONTENT_KEY.value.lower() in str(idx).lower():
|
if "inverted" in str(idx).lower() and Field.CONTENT_KEY.lower() in str(idx).lower():
|
||||||
logger.info("Found existing inverted index: %s", idx)
|
logger.info("Found existing inverted index: %s", idx)
|
||||||
break
|
break
|
||||||
except (RuntimeError, ValueError):
|
except (RuntimeError, ValueError):
|
||||||
@ -669,7 +669,7 @@ class ClickzettaVector(BaseVector):
|
|||||||
|
|
||||||
# Use parameterized INSERT with executemany for better performance and security
|
# Use parameterized INSERT with executemany for better performance and security
|
||||||
# Cast JSON and VECTOR in SQL, pass raw data as parameters
|
# Cast JSON and VECTOR in SQL, pass raw data as parameters
|
||||||
columns = f"id, {Field.CONTENT_KEY.value}, {Field.METADATA_KEY.value}, {Field.VECTOR.value}"
|
columns = f"id, {Field.CONTENT_KEY}, {Field.METADATA_KEY}, {Field.VECTOR}"
|
||||||
insert_sql = (
|
insert_sql = (
|
||||||
f"INSERT INTO {self._config.schema_name}.{self._table_name} ({columns}) "
|
f"INSERT INTO {self._config.schema_name}.{self._table_name} ({columns}) "
|
||||||
f"VALUES (?, ?, CAST(? AS JSON), CAST(? AS VECTOR({vector_dimension})))"
|
f"VALUES (?, ?, CAST(? AS JSON), CAST(? AS VECTOR({vector_dimension})))"
|
||||||
@ -767,7 +767,7 @@ class ClickzettaVector(BaseVector):
|
|||||||
# Use json_extract_string function for ClickZetta compatibility
|
# Use json_extract_string function for ClickZetta compatibility
|
||||||
sql = (
|
sql = (
|
||||||
f"DELETE FROM {self._config.schema_name}.{self._table_name} "
|
f"DELETE FROM {self._config.schema_name}.{self._table_name} "
|
||||||
f"WHERE json_extract_string({Field.METADATA_KEY.value}, '$.{key}') = ?"
|
f"WHERE json_extract_string({Field.METADATA_KEY}, '$.{key}') = ?"
|
||||||
)
|
)
|
||||||
cursor.execute(sql, binding_params=[value])
|
cursor.execute(sql, binding_params=[value])
|
||||||
|
|
||||||
@ -795,9 +795,7 @@ class ClickzettaVector(BaseVector):
|
|||||||
safe_doc_ids = [str(id).replace("'", "''") for id in document_ids_filter]
|
safe_doc_ids = [str(id).replace("'", "''") for id in document_ids_filter]
|
||||||
doc_ids_str = ",".join(f"'{id}'" for id in safe_doc_ids)
|
doc_ids_str = ",".join(f"'{id}'" for id in safe_doc_ids)
|
||||||
# Use json_extract_string function for ClickZetta compatibility
|
# Use json_extract_string function for ClickZetta compatibility
|
||||||
filter_clauses.append(
|
filter_clauses.append(f"json_extract_string({Field.METADATA_KEY}, '$.document_id') IN ({doc_ids_str})")
|
||||||
f"json_extract_string({Field.METADATA_KEY.value}, '$.document_id') IN ({doc_ids_str})"
|
|
||||||
)
|
|
||||||
|
|
||||||
# No need for dataset_id filter since each dataset has its own table
|
# No need for dataset_id filter since each dataset has its own table
|
||||||
|
|
||||||
@ -808,23 +806,21 @@ class ClickzettaVector(BaseVector):
|
|||||||
distance_func = "COSINE_DISTANCE"
|
distance_func = "COSINE_DISTANCE"
|
||||||
if score_threshold > 0:
|
if score_threshold > 0:
|
||||||
query_vector_str = f"CAST('[{self._format_vector_simple(query_vector)}]' AS VECTOR({vector_dimension}))"
|
query_vector_str = f"CAST('[{self._format_vector_simple(query_vector)}]' AS VECTOR({vector_dimension}))"
|
||||||
filter_clauses.append(
|
filter_clauses.append(f"{distance_func}({Field.VECTOR}, {query_vector_str}) < {2 - score_threshold}")
|
||||||
f"{distance_func}({Field.VECTOR.value}, {query_vector_str}) < {2 - score_threshold}"
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
# For L2 distance, smaller is better
|
# For L2 distance, smaller is better
|
||||||
distance_func = "L2_DISTANCE"
|
distance_func = "L2_DISTANCE"
|
||||||
if score_threshold > 0:
|
if score_threshold > 0:
|
||||||
query_vector_str = f"CAST('[{self._format_vector_simple(query_vector)}]' AS VECTOR({vector_dimension}))"
|
query_vector_str = f"CAST('[{self._format_vector_simple(query_vector)}]' AS VECTOR({vector_dimension}))"
|
||||||
filter_clauses.append(f"{distance_func}({Field.VECTOR.value}, {query_vector_str}) < {score_threshold}")
|
filter_clauses.append(f"{distance_func}({Field.VECTOR}, {query_vector_str}) < {score_threshold}")
|
||||||
|
|
||||||
where_clause = " AND ".join(filter_clauses) if filter_clauses else "1=1"
|
where_clause = " AND ".join(filter_clauses) if filter_clauses else "1=1"
|
||||||
|
|
||||||
# Execute vector search query
|
# Execute vector search query
|
||||||
query_vector_str = f"CAST('[{self._format_vector_simple(query_vector)}]' AS VECTOR({vector_dimension}))"
|
query_vector_str = f"CAST('[{self._format_vector_simple(query_vector)}]' AS VECTOR({vector_dimension}))"
|
||||||
search_sql = f"""
|
search_sql = f"""
|
||||||
SELECT id, {Field.CONTENT_KEY.value}, {Field.METADATA_KEY.value},
|
SELECT id, {Field.CONTENT_KEY}, {Field.METADATA_KEY},
|
||||||
{distance_func}({Field.VECTOR.value}, {query_vector_str}) AS distance
|
{distance_func}({Field.VECTOR}, {query_vector_str}) AS distance
|
||||||
FROM {self._config.schema_name}.{self._table_name}
|
FROM {self._config.schema_name}.{self._table_name}
|
||||||
WHERE {where_clause}
|
WHERE {where_clause}
|
||||||
ORDER BY distance
|
ORDER BY distance
|
||||||
@ -887,9 +883,7 @@ class ClickzettaVector(BaseVector):
|
|||||||
safe_doc_ids = [str(id).replace("'", "''") for id in document_ids_filter]
|
safe_doc_ids = [str(id).replace("'", "''") for id in document_ids_filter]
|
||||||
doc_ids_str = ",".join(f"'{id}'" for id in safe_doc_ids)
|
doc_ids_str = ",".join(f"'{id}'" for id in safe_doc_ids)
|
||||||
# Use json_extract_string function for ClickZetta compatibility
|
# Use json_extract_string function for ClickZetta compatibility
|
||||||
filter_clauses.append(
|
filter_clauses.append(f"json_extract_string({Field.METADATA_KEY}, '$.document_id') IN ({doc_ids_str})")
|
||||||
f"json_extract_string({Field.METADATA_KEY.value}, '$.document_id') IN ({doc_ids_str})"
|
|
||||||
)
|
|
||||||
|
|
||||||
# No need for dataset_id filter since each dataset has its own table
|
# No need for dataset_id filter since each dataset has its own table
|
||||||
|
|
||||||
@ -897,13 +891,13 @@ class ClickzettaVector(BaseVector):
|
|||||||
# match_all requires all terms to be present
|
# match_all requires all terms to be present
|
||||||
# Use simple quote escaping for MATCH_ALL since it needs to be in the WHERE clause
|
# Use simple quote escaping for MATCH_ALL since it needs to be in the WHERE clause
|
||||||
escaped_query = query.replace("'", "''")
|
escaped_query = query.replace("'", "''")
|
||||||
filter_clauses.append(f"MATCH_ALL({Field.CONTENT_KEY.value}, '{escaped_query}')")
|
filter_clauses.append(f"MATCH_ALL({Field.CONTENT_KEY}, '{escaped_query}')")
|
||||||
|
|
||||||
where_clause = " AND ".join(filter_clauses)
|
where_clause = " AND ".join(filter_clauses)
|
||||||
|
|
||||||
# Execute full-text search query
|
# Execute full-text search query
|
||||||
search_sql = f"""
|
search_sql = f"""
|
||||||
SELECT id, {Field.CONTENT_KEY.value}, {Field.METADATA_KEY.value}
|
SELECT id, {Field.CONTENT_KEY}, {Field.METADATA_KEY}
|
||||||
FROM {self._config.schema_name}.{self._table_name}
|
FROM {self._config.schema_name}.{self._table_name}
|
||||||
WHERE {where_clause}
|
WHERE {where_clause}
|
||||||
LIMIT {top_k}
|
LIMIT {top_k}
|
||||||
@ -986,19 +980,17 @@ class ClickzettaVector(BaseVector):
|
|||||||
safe_doc_ids = [str(id).replace("'", "''") for id in document_ids_filter]
|
safe_doc_ids = [str(id).replace("'", "''") for id in document_ids_filter]
|
||||||
doc_ids_str = ",".join(f"'{id}'" for id in safe_doc_ids)
|
doc_ids_str = ",".join(f"'{id}'" for id in safe_doc_ids)
|
||||||
# Use json_extract_string function for ClickZetta compatibility
|
# Use json_extract_string function for ClickZetta compatibility
|
||||||
filter_clauses.append(
|
filter_clauses.append(f"json_extract_string({Field.METADATA_KEY}, '$.document_id') IN ({doc_ids_str})")
|
||||||
f"json_extract_string({Field.METADATA_KEY.value}, '$.document_id') IN ({doc_ids_str})"
|
|
||||||
)
|
|
||||||
|
|
||||||
# No need for dataset_id filter since each dataset has its own table
|
# No need for dataset_id filter since each dataset has its own table
|
||||||
|
|
||||||
# Use simple quote escaping for LIKE clause
|
# Use simple quote escaping for LIKE clause
|
||||||
escaped_query = query.replace("'", "''")
|
escaped_query = query.replace("'", "''")
|
||||||
filter_clauses.append(f"{Field.CONTENT_KEY.value} LIKE '%{escaped_query}%'")
|
filter_clauses.append(f"{Field.CONTENT_KEY} LIKE '%{escaped_query}%'")
|
||||||
where_clause = " AND ".join(filter_clauses)
|
where_clause = " AND ".join(filter_clauses)
|
||||||
|
|
||||||
search_sql = f"""
|
search_sql = f"""
|
||||||
SELECT id, {Field.CONTENT_KEY.value}, {Field.METADATA_KEY.value}
|
SELECT id, {Field.CONTENT_KEY}, {Field.METADATA_KEY}
|
||||||
FROM {self._config.schema_name}.{self._table_name}
|
FROM {self._config.schema_name}.{self._table_name}
|
||||||
WHERE {where_clause}
|
WHERE {where_clause}
|
||||||
LIMIT {top_k}
|
LIMIT {top_k}
|
||||||
|
|||||||
@ -57,18 +57,18 @@ class ElasticSearchJaVector(ElasticSearchVector):
|
|||||||
}
|
}
|
||||||
mappings = {
|
mappings = {
|
||||||
"properties": {
|
"properties": {
|
||||||
Field.CONTENT_KEY.value: {
|
Field.CONTENT_KEY: {
|
||||||
"type": "text",
|
"type": "text",
|
||||||
"analyzer": "ja_analyzer",
|
"analyzer": "ja_analyzer",
|
||||||
"search_analyzer": "ja_analyzer",
|
"search_analyzer": "ja_analyzer",
|
||||||
},
|
},
|
||||||
Field.VECTOR.value: { # Make sure the dimension is correct here
|
Field.VECTOR: { # Make sure the dimension is correct here
|
||||||
"type": "dense_vector",
|
"type": "dense_vector",
|
||||||
"dims": dim,
|
"dims": dim,
|
||||||
"index": True,
|
"index": True,
|
||||||
"similarity": "cosine",
|
"similarity": "cosine",
|
||||||
},
|
},
|
||||||
Field.METADATA_KEY.value: {
|
Field.METADATA_KEY: {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
"doc_id": {"type": "keyword"} # Map doc_id to keyword type
|
"doc_id": {"type": "keyword"} # Map doc_id to keyword type
|
||||||
|
|||||||
@ -4,7 +4,7 @@ import math
|
|||||||
from typing import Any, cast
|
from typing import Any, cast
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
import requests
|
from elasticsearch import ConnectionError as ElasticsearchConnectionError
|
||||||
from elasticsearch import Elasticsearch
|
from elasticsearch import Elasticsearch
|
||||||
from flask import current_app
|
from flask import current_app
|
||||||
from packaging.version import parse as parse_version
|
from packaging.version import parse as parse_version
|
||||||
@ -138,7 +138,7 @@ class ElasticSearchVector(BaseVector):
|
|||||||
if not client.ping():
|
if not client.ping():
|
||||||
raise ConnectionError("Failed to connect to Elasticsearch")
|
raise ConnectionError("Failed to connect to Elasticsearch")
|
||||||
|
|
||||||
except requests.ConnectionError as e:
|
except ElasticsearchConnectionError as e:
|
||||||
raise ConnectionError(f"Vector database connection error: {str(e)}")
|
raise ConnectionError(f"Vector database connection error: {str(e)}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise ConnectionError(f"Elasticsearch client initialization failed: {str(e)}")
|
raise ConnectionError(f"Elasticsearch client initialization failed: {str(e)}")
|
||||||
@ -163,9 +163,9 @@ class ElasticSearchVector(BaseVector):
|
|||||||
index=self._collection_name,
|
index=self._collection_name,
|
||||||
id=uuids[i],
|
id=uuids[i],
|
||||||
document={
|
document={
|
||||||
Field.CONTENT_KEY.value: documents[i].page_content,
|
Field.CONTENT_KEY: documents[i].page_content,
|
||||||
Field.VECTOR.value: embeddings[i] or None,
|
Field.VECTOR: embeddings[i] or None,
|
||||||
Field.METADATA_KEY.value: documents[i].metadata or {},
|
Field.METADATA_KEY: documents[i].metadata or {},
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
self._client.indices.refresh(index=self._collection_name)
|
self._client.indices.refresh(index=self._collection_name)
|
||||||
@ -193,7 +193,7 @@ class ElasticSearchVector(BaseVector):
|
|||||||
def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:
|
def search_by_vector(self, query_vector: list[float], **kwargs: Any) -> list[Document]:
|
||||||
top_k = kwargs.get("top_k", 4)
|
top_k = kwargs.get("top_k", 4)
|
||||||
num_candidates = math.ceil(top_k * 1.5)
|
num_candidates = math.ceil(top_k * 1.5)
|
||||||
knn = {"field": Field.VECTOR.value, "query_vector": query_vector, "k": top_k, "num_candidates": num_candidates}
|
knn = {"field": Field.VECTOR, "query_vector": query_vector, "k": top_k, "num_candidates": num_candidates}
|
||||||
document_ids_filter = kwargs.get("document_ids_filter")
|
document_ids_filter = kwargs.get("document_ids_filter")
|
||||||
if document_ids_filter:
|
if document_ids_filter:
|
||||||
knn["filter"] = {"terms": {"metadata.document_id": document_ids_filter}}
|
knn["filter"] = {"terms": {"metadata.document_id": document_ids_filter}}
|
||||||
@ -205,9 +205,9 @@ class ElasticSearchVector(BaseVector):
|
|||||||
docs_and_scores.append(
|
docs_and_scores.append(
|
||||||
(
|
(
|
||||||
Document(
|
Document(
|
||||||
page_content=hit["_source"][Field.CONTENT_KEY.value],
|
page_content=hit["_source"][Field.CONTENT_KEY],
|
||||||
vector=hit["_source"][Field.VECTOR.value],
|
vector=hit["_source"][Field.VECTOR],
|
||||||
metadata=hit["_source"][Field.METADATA_KEY.value],
|
metadata=hit["_source"][Field.METADATA_KEY],
|
||||||
),
|
),
|
||||||
hit["_score"],
|
hit["_score"],
|
||||||
)
|
)
|
||||||
@ -224,13 +224,13 @@ class ElasticSearchVector(BaseVector):
|
|||||||
return docs
|
return docs
|
||||||
|
|
||||||
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
||||||
query_str: dict[str, Any] = {"match": {Field.CONTENT_KEY.value: query}}
|
query_str: dict[str, Any] = {"match": {Field.CONTENT_KEY: query}}
|
||||||
document_ids_filter = kwargs.get("document_ids_filter")
|
document_ids_filter = kwargs.get("document_ids_filter")
|
||||||
|
|
||||||
if document_ids_filter:
|
if document_ids_filter:
|
||||||
query_str = {
|
query_str = {
|
||||||
"bool": {
|
"bool": {
|
||||||
"must": {"match": {Field.CONTENT_KEY.value: query}},
|
"must": {"match": {Field.CONTENT_KEY: query}},
|
||||||
"filter": {"terms": {"metadata.document_id": document_ids_filter}},
|
"filter": {"terms": {"metadata.document_id": document_ids_filter}},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@ -240,9 +240,9 @@ class ElasticSearchVector(BaseVector):
|
|||||||
for hit in results["hits"]["hits"]:
|
for hit in results["hits"]["hits"]:
|
||||||
docs.append(
|
docs.append(
|
||||||
Document(
|
Document(
|
||||||
page_content=hit["_source"][Field.CONTENT_KEY.value],
|
page_content=hit["_source"][Field.CONTENT_KEY],
|
||||||
vector=hit["_source"][Field.VECTOR.value],
|
vector=hit["_source"][Field.VECTOR],
|
||||||
metadata=hit["_source"][Field.METADATA_KEY.value],
|
metadata=hit["_source"][Field.METADATA_KEY],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -270,14 +270,14 @@ class ElasticSearchVector(BaseVector):
|
|||||||
dim = len(embeddings[0])
|
dim = len(embeddings[0])
|
||||||
mappings = {
|
mappings = {
|
||||||
"properties": {
|
"properties": {
|
||||||
Field.CONTENT_KEY.value: {"type": "text"},
|
Field.CONTENT_KEY: {"type": "text"},
|
||||||
Field.VECTOR.value: { # Make sure the dimension is correct here
|
Field.VECTOR: { # Make sure the dimension is correct here
|
||||||
"type": "dense_vector",
|
"type": "dense_vector",
|
||||||
"dims": dim,
|
"dims": dim,
|
||||||
"index": True,
|
"index": True,
|
||||||
"similarity": "cosine",
|
"similarity": "cosine",
|
||||||
},
|
},
|
||||||
Field.METADATA_KEY.value: {
|
Field.METADATA_KEY: {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
"doc_id": {"type": "keyword"}, # Map doc_id to keyword type
|
"doc_id": {"type": "keyword"}, # Map doc_id to keyword type
|
||||||
|
|||||||
@ -67,9 +67,9 @@ class HuaweiCloudVector(BaseVector):
|
|||||||
index=self._collection_name,
|
index=self._collection_name,
|
||||||
id=uuids[i],
|
id=uuids[i],
|
||||||
document={
|
document={
|
||||||
Field.CONTENT_KEY.value: documents[i].page_content,
|
Field.CONTENT_KEY: documents[i].page_content,
|
||||||
Field.VECTOR.value: embeddings[i] or None,
|
Field.VECTOR: embeddings[i] or None,
|
||||||
Field.METADATA_KEY.value: documents[i].metadata or {},
|
Field.METADATA_KEY: documents[i].metadata or {},
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
self._client.indices.refresh(index=self._collection_name)
|
self._client.indices.refresh(index=self._collection_name)
|
||||||
@ -101,7 +101,7 @@ class HuaweiCloudVector(BaseVector):
|
|||||||
"size": top_k,
|
"size": top_k,
|
||||||
"query": {
|
"query": {
|
||||||
"vector": {
|
"vector": {
|
||||||
Field.VECTOR.value: {
|
Field.VECTOR: {
|
||||||
"vector": query_vector,
|
"vector": query_vector,
|
||||||
"topk": top_k,
|
"topk": top_k,
|
||||||
}
|
}
|
||||||
@ -116,9 +116,9 @@ class HuaweiCloudVector(BaseVector):
|
|||||||
docs_and_scores.append(
|
docs_and_scores.append(
|
||||||
(
|
(
|
||||||
Document(
|
Document(
|
||||||
page_content=hit["_source"][Field.CONTENT_KEY.value],
|
page_content=hit["_source"][Field.CONTENT_KEY],
|
||||||
vector=hit["_source"][Field.VECTOR.value],
|
vector=hit["_source"][Field.VECTOR],
|
||||||
metadata=hit["_source"][Field.METADATA_KEY.value],
|
metadata=hit["_source"][Field.METADATA_KEY],
|
||||||
),
|
),
|
||||||
hit["_score"],
|
hit["_score"],
|
||||||
)
|
)
|
||||||
@ -135,15 +135,15 @@ class HuaweiCloudVector(BaseVector):
|
|||||||
return docs
|
return docs
|
||||||
|
|
||||||
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
def search_by_full_text(self, query: str, **kwargs: Any) -> list[Document]:
|
||||||
query_str = {"match": {Field.CONTENT_KEY.value: query}}
|
query_str = {"match": {Field.CONTENT_KEY: query}}
|
||||||
results = self._client.search(index=self._collection_name, query=query_str, size=kwargs.get("top_k", 4))
|
results = self._client.search(index=self._collection_name, query=query_str, size=kwargs.get("top_k", 4))
|
||||||
docs = []
|
docs = []
|
||||||
for hit in results["hits"]["hits"]:
|
for hit in results["hits"]["hits"]:
|
||||||
docs.append(
|
docs.append(
|
||||||
Document(
|
Document(
|
||||||
page_content=hit["_source"][Field.CONTENT_KEY.value],
|
page_content=hit["_source"][Field.CONTENT_KEY],
|
||||||
vector=hit["_source"][Field.VECTOR.value],
|
vector=hit["_source"][Field.VECTOR],
|
||||||
metadata=hit["_source"][Field.METADATA_KEY.value],
|
metadata=hit["_source"][Field.METADATA_KEY],
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@ -171,8 +171,8 @@ class HuaweiCloudVector(BaseVector):
|
|||||||
dim = len(embeddings[0])
|
dim = len(embeddings[0])
|
||||||
mappings = {
|
mappings = {
|
||||||
"properties": {
|
"properties": {
|
||||||
Field.CONTENT_KEY.value: {"type": "text"},
|
Field.CONTENT_KEY: {"type": "text"},
|
||||||
Field.VECTOR.value: { # Make sure the dimension is correct here
|
Field.VECTOR: { # Make sure the dimension is correct here
|
||||||
"type": "vector",
|
"type": "vector",
|
||||||
"dimension": dim,
|
"dimension": dim,
|
||||||
"indexing": True,
|
"indexing": True,
|
||||||
@ -181,7 +181,7 @@ class HuaweiCloudVector(BaseVector):
|
|||||||
"neighbors": 32,
|
"neighbors": 32,
|
||||||
"efc": 128,
|
"efc": 128,
|
||||||
},
|
},
|
||||||
Field.METADATA_KEY.value: {
|
Field.METADATA_KEY: {
|
||||||
"type": "object",
|
"type": "object",
|
||||||
"properties": {
|
"properties": {
|
||||||
"doc_id": {"type": "keyword"} # Map doc_id to keyword type
|
"doc_id": {"type": "keyword"} # Map doc_id to keyword type
|
||||||
|
|||||||
@ -125,9 +125,9 @@ class LindormVectorStore(BaseVector):
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
action_values: dict[str, Any] = {
|
action_values: dict[str, Any] = {
|
||||||
Field.CONTENT_KEY.value: documents[i].page_content,
|
Field.CONTENT_KEY: documents[i].page_content,
|
||||||
Field.VECTOR.value: embeddings[i],
|
Field.VECTOR: embeddings[i],
|
||||||
Field.METADATA_KEY.value: documents[i].metadata,
|
Field.METADATA_KEY: documents[i].metadata,
|
||||||
}
|
}
|
||||||
if self._using_ugc:
|
if self._using_ugc:
|
||||||
action_header["index"]["routing"] = self._routing
|
action_header["index"]["routing"] = self._routing
|
||||||
@ -149,7 +149,7 @@ class LindormVectorStore(BaseVector):
|
|||||||
|
|
||||||
def get_ids_by_metadata_field(self, key: str, value: str):
|
def get_ids_by_metadata_field(self, key: str, value: str):
|
||||||
query: dict[str, Any] = {
|
query: dict[str, Any] = {
|
||||||
"query": {"bool": {"must": [{"term": {f"{Field.METADATA_KEY.value}.{key}.keyword": value}}]}}
|
"query": {"bool": {"must": [{"term": {f"{Field.METADATA_KEY}.{key}.keyword": value}}]}}
|
||||||
}
|
}
|
||||||
if self._using_ugc:
|
if self._using_ugc:
|
||||||
query["query"]["bool"]["must"].append({"term": {f"{ROUTING_FIELD}.keyword": self._routing}})
|
query["query"]["bool"]["must"].append({"term": {f"{ROUTING_FIELD}.keyword": self._routing}})
|
||||||
@ -252,14 +252,14 @@ class LindormVectorStore(BaseVector):
|
|||||||
search_query: dict[str, Any] = {
|
search_query: dict[str, Any] = {
|
||||||
"size": top_k,
|
"size": top_k,
|
||||||
"_source": True,
|
"_source": True,
|
||||||
"query": {"knn": {Field.VECTOR.value: {"vector": query_vector, "k": top_k}}},
|
"query": {"knn": {Field.VECTOR: {"vector": query_vector, "k": top_k}}},
|
||||||
}
|
}
|
||||||
|
|
||||||
final_ext: dict[str, Any] = {"lvector": {}}
|
final_ext: dict[str, Any] = {"lvector": {}}
|
||||||
if filters is not None and len(filters) > 0:
|
if filters is not None and len(filters) > 0:
|
||||||
# when using filter, transform filter from List[Dict] to Dict as valid format
|
# when using filter, transform filter from List[Dict] to Dict as valid format
|
||||||
filter_dict = {"bool": {"must": filters}} if len(filters) > 1 else filters[0]
|
filter_dict = {"bool": {"must": filters}} if len(filters) > 1 else filters[0]
|
||||||
search_query["query"]["knn"][Field.VECTOR.value]["filter"] = filter_dict # filter should be Dict
|
search_query["query"]["knn"][Field.VECTOR]["filter"] = filter_dict # filter should be Dict
|
||||||
final_ext["lvector"]["filter_type"] = "pre_filter"
|
final_ext["lvector"]["filter_type"] = "pre_filter"
|
||||||
|
|
||||||
if final_ext != {"lvector": {}}:
|
if final_ext != {"lvector": {}}:
|
||||||
@ -279,9 +279,9 @@ class LindormVectorStore(BaseVector):
|
|||||||
docs_and_scores.append(
|
docs_and_scores.append(
|
||||||
(
|
(
|
||||||
Document(
|
Document(
|
||||||
page_content=hit["_source"][Field.CONTENT_KEY.value],
|
page_content=hit["_source"][Field.CONTENT_KEY],
|
||||||
vector=hit["_source"][Field.VECTOR.value],
|
vector=hit["_source"][Field.VECTOR],
|
||||||
metadata=hit["_source"][Field.METADATA_KEY.value],
|
metadata=hit["_source"][Field.METADATA_KEY],
|
||||||
),
|
),
|
||||||
hit["_score"],
|
hit["_score"],
|
||||||
)
|
)
|
||||||
@ -318,9 +318,9 @@ class LindormVectorStore(BaseVector):
|
|||||||
|
|
||||||
docs = []
|
docs = []
|
||||||
for hit in response["hits"]["hits"]:
|
for hit in response["hits"]["hits"]:
|
||||||
metadata = hit["_source"].get(Field.METADATA_KEY.value)
|
metadata = hit["_source"].get(Field.METADATA_KEY)
|
||||||
vector = hit["_source"].get(Field.VECTOR.value)
|
vector = hit["_source"].get(Field.VECTOR)
|
||||||
page_content = hit["_source"].get(Field.CONTENT_KEY.value)
|
page_content = hit["_source"].get(Field.CONTENT_KEY)
|
||||||
doc = Document(page_content=page_content, vector=vector, metadata=metadata)
|
doc = Document(page_content=page_content, vector=vector, metadata=metadata)
|
||||||
docs.append(doc)
|
docs.append(doc)
|
||||||
|
|
||||||
@ -342,8 +342,8 @@ class LindormVectorStore(BaseVector):
|
|||||||
"settings": {"index": {"knn": True, "knn_routing": self._using_ugc}},
|
"settings": {"index": {"knn": True, "knn_routing": self._using_ugc}},
|
||||||
"mappings": {
|
"mappings": {
|
||||||
"properties": {
|
"properties": {
|
||||||
Field.CONTENT_KEY.value: {"type": "text"},
|
Field.CONTENT_KEY: {"type": "text"},
|
||||||
Field.VECTOR.value: {
|
Field.VECTOR: {
|
||||||
"type": "knn_vector",
|
"type": "knn_vector",
|
||||||
"dimension": len(embeddings[0]), # Make sure the dimension is correct here
|
"dimension": len(embeddings[0]), # Make sure the dimension is correct here
|
||||||
"method": {
|
"method": {
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Loading…
Reference in New Issue
Block a user