Skip to content
51 changes: 43 additions & 8 deletions packages/dbgpt-app/src/dbgpt_app/knowledge/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,10 @@ async def space_add(


@router.post("/knowledge/space/list")
async def space_list(request: KnowledgeSpaceRequest):
async def space_list(
request: KnowledgeSpaceRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/space/list params: {request}")
try:
res = await blocking_func_to_async(
Expand All @@ -110,7 +113,10 @@ async def space_list(request: KnowledgeSpaceRequest):


@router.post("/knowledge/space/delete")
def space_delete(request: KnowledgeSpaceRequest):
def space_delete(
request: KnowledgeSpaceRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/space/delete params: {request}")
try:
# delete Files in 'pilot/data/
Expand Down Expand Up @@ -144,7 +150,10 @@ async def retrieve_strategy_list(


@router.post("/knowledge/{space_id}/arguments")
async def arguments(space_id: str):
async def arguments(
space_id: str,
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/knowledge/{space_id}/arguments params: {space_id}")
try:
res = await blocking_func_to_async(
Expand Down Expand Up @@ -173,6 +182,7 @@ async def recall_test(
@router.get("/knowledge/{space_id}/recall_retrievers")
def recall_retrievers(
space_id: str,
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/knowledge/{space_id}/recall_retrievers params:")
try:
Expand Down Expand Up @@ -230,7 +240,11 @@ async def arguments_save(


@router.post("/knowledge/{space_name}/document/add")
async def document_add(space_name: str, request: KnowledgeDocumentRequest):
async def document_add(
space_name: str,
request: KnowledgeDocumentRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/document/add params: {space_name}, {request}")
try:
res = await blocking_func_to_async(
Expand All @@ -250,6 +264,7 @@ def document_edit(
space_name: str,
request: KnowledgeDocumentRequest,
service: Service = Depends(get_rag_service),
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/document/edit params: {space_name}, {request}")
space = service.get({"name": space_name})
Expand Down Expand Up @@ -369,7 +384,11 @@ def document_list(


@router.post("/knowledge/{space_name}/graphvis")
def graph_vis(space_name: str, query_request: GraphVisRequest):
def graph_vis(
space_name: str,
query_request: GraphVisRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"/document/list params: {space_name}, {query_request}")
try:
return Result.succ(
Expand All @@ -382,7 +401,11 @@ def graph_vis(space_name: str, query_request: GraphVisRequest):


@router.post("/knowledge/{space_name}/document/delete")
def document_delete(space_name: str, query_request: DocumentQueryRequest):
def document_delete(
space_name: str,
query_request: DocumentQueryRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
print(f"/document/list params: {space_name}, {query_request}")
try:
return Result.succ(
Expand All @@ -399,6 +422,7 @@ async def document_upload(
doc_type: str = Form(...),
doc_file: UploadFile = File(...),
fs: FileStorageClient = Depends(get_fs),
user_token: UserRequest = Depends(get_user_from_headers),
):
print(f"/document/upload params: {space_name}")
try:
Expand Down Expand Up @@ -467,6 +491,7 @@ async def document_sync(
space_name: str,
request: DocumentSyncRequest,
service: Service = Depends(get_rag_service),
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"Received params: {space_name}, {request}")
try:
Expand Down Expand Up @@ -496,6 +521,7 @@ async def batch_document_sync(
space_name: str,
request: List[KnowledgeSyncRequest],
service: Service = Depends(get_rag_service),
user_token: UserRequest = Depends(get_user_from_headers),
):
logger.info(f"Received params: {space_name}, {request}")
try:
Expand All @@ -517,6 +543,7 @@ def chunk_list(
space_name: str,
query_request: ChunkQueryRequest,
service: Service = Depends(get_rag_service),
user_token: UserRequest = Depends(get_user_from_headers),
):
print(f"/chunk/list params: {space_name}, {query_request}")
try:
Expand Down Expand Up @@ -545,6 +572,7 @@ def chunk_edit(
space_name: str,
edit_request: ChunkEditRequest,
service: Service = Depends(get_rag_service),
user_token: UserRequest = Depends(get_user_from_headers),
):
print(f"/chunk/edit params: {space_name}, {edit_request}")
try:
Expand All @@ -556,7 +584,11 @@ def chunk_edit(


@router.post("/knowledge/{vector_name}/query")
def similarity_query(space_name: str, query_request: KnowledgeQueryRequest):
def similarity_query(
space_name: str,
query_request: KnowledgeQueryRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
print(f"Received params: {space_name}, {query_request}")
storage_manager = StorageManager.get_instance(CFG.SYSTEM_APP)
vector_store_connector = storage_manager.create_vector_store(index_name=space_name)
Expand All @@ -572,7 +604,10 @@ def similarity_query(space_name: str, query_request: KnowledgeQueryRequest):


@router.post("/knowledge/document/summary")
async def document_summary(request: DocumentSummaryRequest):
async def document_summary(
request: DocumentSummaryRequest,
user_token: UserRequest = Depends(get_user_from_headers),
):
print(f"/document/summary params: {request}")
try:
with root_tracer.start_span(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3992,6 +3992,7 @@ async def delete_share_link(
@router.get("/v1/agent/files/download")
async def download_agent_file(
file_path: str = Query(..., description="Absolute path to the file to download"),
user_token: UserRequest = Depends(get_user_from_headers),
):
"""Download a file created by agent tools (shell_interpreter, code_interpreter).

Expand All @@ -4001,11 +4002,11 @@ async def download_agent_file(
from fastapi import HTTPException
from fastapi.responses import FileResponse

from dbgpt.configs.model_config import PILOT_PATH, ROOT_PATH
from dbgpt.configs.model_config import PILOT_PATH

# If path is not absolute, resolve relative to ROOT_PATH (sandbox working dir)
# If path is not absolute, resolve it under the controlled agent temp directory.
if not os.path.isabs(file_path):
file_path = os.path.join(ROOT_PATH, file_path)
file_path = os.path.join(PILOT_PATH, "tmp", file_path)

# Resolve to absolute path and prevent path traversal
try:
Expand All @@ -4017,7 +4018,6 @@ async def download_agent_file(
allowed_dirs = [
os.path.realpath("/tmp"),
os.path.realpath(os.path.join(PILOT_PATH, "tmp")),
os.path.realpath(ROOT_PATH),
]

if not any(resolved.startswith(d + os.sep) or resolved == d for d in allowed_dirs):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,36 @@
logger = logging.getLogger(__name__)


def _strip_sql_comments(sql: str) -> str:
result = []
index = 0
length = len(sql)

while index < length:
if sql.startswith("/*", index):
end = sql.find("*/", index + 2)
if end == -1:
result.append(" ")
break
result.append(" ")
index = end + 2
continue

if sql.startswith("--", index):
end = sql.find("\n", index + 2)
if end == -1:
result.append(" ")
break
result.append(" ")
index = end
continue

result.append(sql[index])
index += 1

return "".join(result)


def get_conversation_serve() -> ConversationServe:
return ConversationServe.get_instance(CFG.SYSTEM_APP)

Expand Down Expand Up @@ -102,8 +132,7 @@ def sanitize_sql(sql: str, db_type: str = None) -> Tuple[bool, str, dict]:
Tuple of (is_safe, reason, params)
"""
# Normalize SQL (remove comments and excess whitespace)
sql = re.sub(r"/\*.*?\*/", " ", sql)
sql = re.sub(r"--.*?$", " ", sql, flags=re.MULTILINE)
sql = _strip_sql_comments(sql)
sql = re.sub(r"\s+", " ", sql).strip()

# Block multiple statements
Expand Down Expand Up @@ -191,7 +220,7 @@ async def editor_sql_run(run_param: dict = Body()):
# Sanitize and parameterize the SQL query
is_safe, result, params = sanitize_sql(sql, db_type)
if not is_safe:
logger.warning(f"Blocked dangerous SQL: {sql}")
logger.warning("Blocked dangerous SQL with length %s", len(sql))
return Result.failed(msg=f"Operation not allowed: {result}")

try:
Expand Down Expand Up @@ -272,7 +301,7 @@ async def chart_run(run_param: dict = Body()):
# Sanitize and parameterize the SQL query
is_safe, result, params = sanitize_sql(sql, db_type)
if not is_safe:
logger.warning(f"Blocked dangerous SQL: {sql}")
logger.warning("Blocked dangerous SQL with length %s", len(sql))
return Result.failed(msg=f"Operation not allowed: {result}")

try:
Expand Down
25 changes: 20 additions & 5 deletions packages/dbgpt-app/src/dbgpt_app/openapi/api_v2.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,4 @@
import json
import re
import time
import uuid
from typing import AsyncIterator, Optional
Expand Down Expand Up @@ -232,6 +231,24 @@ async def no_stream_wrapper(
)


def _extract_sse_json_payload(output: str) -> Optional[dict]:
data_index = output.find("data:")
if data_index < 0:
return None

payload = output[data_index + len("data:") :].lstrip()
if not payload or payload[0] != "{":
return None

try:
parsed, _ = json.JSONDecoder().raw_decode(payload)
except json.JSONDecodeError:
return None
if not isinstance(parsed, dict):
return None
return parsed


async def chat_app_stream_wrapper(request: ChatCompletionRequestBody = None):
"""chat app stream
Args:
Expand All @@ -245,10 +262,8 @@ async def chat_app_stream_wrapper(request: ChatCompletionRequestBody = None):
user_code=request.user_name,
sys_code=request.sys_code,
):
match = re.search(r"data:\s*({.*})", output)
if match:
json_str = match.group(1)
vis = json.loads(json_str)
vis = _extract_sse_json_payload(output)
if vis:
vis_content = vis.get("vis", None)
if vis_content != "[DONE]":
choice_data = ChatCompletionResponseStreamChoice(
Expand Down
Loading
Loading