FastApi-Swagger
一、FastApi中怎么使用Swagger连接数据库调用大模型返回sql
1. 基本设置
首先安装必要的依赖:
pip install fastapi uvicorn sqlalchemy openai python-dotenv
2. 创建FastAPI应用并集成Swagger
from fastapi import FastAPI, HTTPException from pydantic import BaseModel from typing import Optional import os from sqlalchemy import create_engine, text from sqlalchemy.exc import SQLAlchemyError import openai app = FastAPI() # 配置Swagger文档 app.title = "Database SQL Generator API" app.description = "API for generating SQL queries using AI models" app.version = "1.0.0"
3. 定义请求模型
class DatabaseConnection(BaseModel):
db_type: str # e.g., "mysql", "postgresql", "sqlite"
host: str
port: int
username: str
password: str
database: str
ssl: Optional[bool] = False
class QueryRequest(BaseModel):
connection: DatabaseConnection
natural_language_query: str
model: Optional[str] = "gpt-3.5-turbo" # 或其他大模型
temperature: Optional[float] = 0.3
4. 实现核心功能
def get_db_connection_string(conn: DatabaseConnection) -> str:
if conn.db_type == "mysql":
return f"mysql+pymysql://{conn.username}:{conn.password}@{conn.host}:{conn.port}/{conn.database}"
elif conn.db_type == "postgresql":
return f"postgresql://{conn.username}:{conn.password}@{conn.host}:{conn.port}/{conn.database}"
else:
raise HTTPException(status_code=400, detail="Unsupported database type")
def get_database_schema(conn_str: str) -> str:
try:
engine = create_engine(conn_str)
with engine.connect() as connection:
# 获取表信息
tables = connection.execute(text("""
SELECT table_name, column_name, data_type
FROM information_schema.columns
WHERE table_schema = 'public'
ORDER BY table_name, ordinal_position
"""))
schema = {}
for table in tables:
if table.table_name not in schema:
schema[table.table_name] = []
schema[table.table_name].append(f"{table.column_name} ({table.data_type})")
return "\n".join([f"Table {table}:\n " + "\n ".join(columns)
for table, columns in schema.items()])
except SQLAlchemyError as e:
raise HTTPException(status_code=500, detail=f"Database connection error: {str(e)}")
def generate_sql_with_ai(schema: str, query: str, model: str = "gpt-3.5-turbo") -> str:
try:
prompt = f"""
Given the following database schema:
{schema}
Write a SQL query to: {query}
Return only the SQL query without any additional explanation or formatting.
"""
response = openai.ChatCompletion.create(
model=model,
messages=[
{"role": "system", "content": "You are a helpful SQL assistant."},
{"role": "user", "content": prompt}
],
temperature=0.3
)
return response.choices[0].message.content.strip()
except Exception as e:
raise HTTPException(status_code=500, detail=f"AI model error: {str(e)}")
5. 创建API端点
@app.post("/generate-sql", summary="Generate SQL from natural language")
async def generate_sql(request: QueryRequest):
"""
Generate SQL query from natural language using AI model.
- **connection**: Database connection details
- **natural_language_query**: Your query in natural language
- **model**: AI model to use (default: gpt-3.5-turbo)
- **temperature**: Creativity of the model (0-1)
"""
# 1. 获取数据库连接字符串
conn_str = get_db_connection_string(request.connection)
# 2. 获取数据库模式
schema = get_database_schema(conn_str)
# 3. 使用大模型生成SQL
sql_query = generate_sql_with_ai(
schema=schema,
query=request.natural_language_query,
model=request.model
)
return {
"status": "success",
"sql_query": sql_query,
"database_schema": schema
}
6. 运行应用
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)
7. 环境变量配置
创建.env文件:
OPENAI_API_KEY=your_openai_api_key
8. 使用说明
-
启动FastAPI应用:
uvicorn main:app --reload -
访问Swagger UI:
http://localhost:8000/docs -
在Swagger UI中:
-
找到
/generate-sql端点 -
点击"Try it out"
-
填写数据库连接信息和自然语言查询
-
执行请求
-
二、FastApi中怎么使用Swagger连接数据库调用本地或者他人的大模型返回sql
1. 安装依赖
pip install fastapi uvicorn sqlalchemy python-dotenv requests # 根据选择的模型可能需要额外安装: # pip install transformers torch # 本地模型 # pip install openai # OpenAI API
2. 创建FastAPI应用
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from typing import Optional, Literal
import os
from sqlalchemy import create_engine, text
from sqlalchemy.exc import SQLAlchemyError
import requests
from dotenv import load_dotenv
load_dotenv()
app = FastAPI(
title="SQL Generator API",
description="Generate SQL queries using local or cloud-based LLMs",
version="1.0.0",
openapi_tags=[{
"name": "SQL Generation",
"description": "Endpoints for generating SQL from natural language"
}]
)
3. 定义数据模型
class DatabaseConnection(BaseModel):
db_type: Literal["mysql", "postgresql", "sqlite", "mssql"]
host: str
port: int
username: str
password: str
database: str
ssl: Optional[bool] = False
class ModelConfig(BaseModel):
model_type: Literal["local", "openai", "anthropic", "deepseek", "custom"]
model_name: str # 如 "llama2-7b", "gpt-3.5-turbo"等
api_base: Optional[str] = None # 自定义API基础URL
api_key: Optional[str] = None
class QueryRequest(BaseModel):
connection: DatabaseConnection
natural_language_query: str
model_config: ModelConfig
temperature: Optional[float] = 0.3
max_tokens: Optional[int] = 500
4. 数据库工具函数
def get_db_connection_string(conn: DatabaseConnection) -> str:
"""生成数据库连接字符串"""
drivers = {
"mysql": "mysql+pymysql",
"postgresql": "postgresql",
"sqlite": "sqlite",
"mssql": "mssql+pyodbc"
}
if conn.db_type not in drivers:
raise HTTPException(status_code=400, detail="Unsupported database type")
if conn.db_type == "sqlite":
return f"sqlite:///{conn.database}"
return f"{drivers[conn.db_type]}://{conn.username}:{conn.password}@{conn.host}:{conn.port}/{conn.database}"
def get_database_schema(conn_str: str, db_type: str) -> str:
"""获取数据库模式信息"""
try:
engine = create_engine(conn_str)
with engine.connect() as connection:
if db_type == "mysql":
query = """
SELECT table_name, column_name, data_type
FROM information_schema.columns
WHERE table_schema = DATABASE()
ORDER BY table_name, ordinal_position
"""
elif db_type == "postgresql":
query = """
SELECT table_name, column_name, data_type
FROM information_schema.columns
WHERE table_schema = 'public'
ORDER BY table_name, ordinal_position
"""
else:
# 其他数据库类型的查询
query = """
SELECT table_name, column_name, data_type
FROM information_schema.columns
ORDER BY table_name, ordinal_position
"""
tables = connection.execute(text(query))
schema = {}
for table in tables:
if table.table_name not in schema:
schema[table.table_name] = []
schema[table.table_name].append(f"{table.column_name} ({table.data_type})")
return "\n".join([f"Table {table}:\n " + "\n ".join(columns)
for table, columns in schema.items()])
except SQLAlchemyError as e:
raise HTTPException(status_code=500, detail=f"Database connection error: {str(e)}")
5. 大模型调用函数
def call_model(prompt: str, model_config: ModelConfig, temperature: float = 0.3, max_tokens: int = 500) -> str:
"""调用不同的大模型API"""
try:
if model_config.model_type == "local":
# 本地模型调用 - 需要根据实际部署方式调整
return call_local_model(prompt, model_config)
elif model_config.model_type == "openai":
return call_openai_api(prompt, model_config, temperature, max_tokens)
elif model_config.model_type == "deepseek":
return call_deepseek_api(prompt, model_config, temperature, max_tokens)
elif model_config.model_type == "custom":
return call_custom_api(prompt, model_config, temperature, max_tokens)
else:
raise HTTPException(status_code=400, detail="Unsupported model type")
except Exception as e:
raise HTTPException(status_code=500, detail=f"Model API error: {str(e)}")
def call_local_model(prompt: str, model_config: ModelConfig) -> str:
"""调用本地部署的模型"""
# 示例:使用transformers库调用本地模型
# 实际使用时需要根据模型部署方式调整
# 方法1:使用HTTP API(如果模型以API方式部署)
if model_config.api_base:
response = requests.post(
f"{model_config.api_base}/generate",
json={
"prompt": prompt,
"model": model_config.model_name,
"max_tokens": 500
},
headers={"Authorization": f"Bearer {model_config.api_key}"} if model_config.api_key else {}
)
return response.json()["text"]
# 方法2:直接加载模型(不推荐在生产环境使用)
from transformers import pipeline
generator = pipeline('text-generation', model=model_config.model_name)
result = generator(prompt, max_length=500)
return result[0]['generated_text']
def call_openai_api(prompt: str, model_config: ModelConfig, temperature: float, max_tokens: int) -> str:
"""调用OpenAI API"""
import openai
openai.api_key = model_config.api_key or os.getenv("OPENAI_API_KEY")
response = openai.ChatCompletion.create(
model=model_config.model_name,
messages=[
{"role": "system", "content": "You are a SQL expert. Generate only SQL code without explanations."},
{"role": "user", "content": prompt}
],
temperature=temperature,
max_tokens=max_tokens
)
return response.choices[0].message.content
def call_deepseek_api(prompt: str, model_config: ModelConfig, temperature: float, max_tokens: int) -> str:
"""调用DeepSeek API"""
headers = {
"Authorization": f"Bearer {model_config.api_key or os.getenv('DEEPSEEK_API_KEY')}",
"Content-Type": "application/json"
}
data = {
"model": model_config.model_name,
"messages": [
{"role": "system", "content": "You are a SQL expert. Return only SQL code."},
{"role": "user", "content": prompt}
],
"temperature": temperature,
"max_tokens": max_tokens
}
response = requests.post(
model_config.api_base or "https://api.deepseek.com/v1/chat/completions",
headers=headers,
json=data
)
return response.json()["choices"][0]["message"]["content"]
def call_custom_api(prompt: str, model_config: ModelConfig, temperature: float, max_tokens: int) -> str:
"""调用自定义API"""
if not model_config.api_base:
raise HTTPException(status_code=400, detail="Custom API requires api_base")
response = requests.post(
model_config.api_base,
json={
"prompt": prompt,
"model": model_config.model_name,
"temperature": temperature,
"max_tokens": max_tokens
},
headers={"Authorization": f"Bearer {model_config.api_key}"} if model_config.api_key else {}
)
return response.json()["text"]
6. 创建API端点
@app.post("/generate-sql",
tags=["SQL Generation"],
summary="Generate SQL from natural language",
response_description="The generated SQL query")
async def generate_sql(request: QueryRequest):
"""
Generate SQL query from natural language using specified LLM.
- **connection**: Database connection details
- **natural_language_query**: Query in natural language
- **model_config**: Configuration for the LLM to use
- **temperature**: Creativity parameter (0-1)
- **max_tokens**: Maximum length of the response
"""
# 1. 获取数据库连接字符串
conn_str = get_db_connection_string(request.connection)
# 2. 获取数据库模式
schema = get_database_schema(conn_str, request.connection.db_type)
# 3. 构建提示词
prompt = f"""
Database schema:
{schema}
Task: Convert the following natural language query to SQL:
"{request.natural_language_query}"
Requirements:
- Return only the SQL query
- Do not include any explanations or additional text
- Use proper SQL syntax for {request.connection.db_type}
"""
# 4. 调用大模型生成SQL
sql_query = call_model(
prompt=prompt,
model_config=request.model_config,
temperature=request.temperature,
max_tokens=request.max_tokens
)
# 5. 清理响应(去除可能的额外文本)
sql_query = sql_query.strip()
if sql_query.startswith("```sql"):
sql_query = sql_query[6:-3].strip()
elif sql_query.startswith("```"):
sql_query = sql_query[3:-3].strip()
return {
"status": "success",
"sql_query": sql_query,
"database_schema": schema,
"model_used": request.model_config.model_name
}
7. 运行应用
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000)
使用说明
-
本地模型部署:
-
使用Transformers或类似库部署本地模型
-
可以通过API方式暴露模型(如使用FastAPI)
-
在
ModelConfig中设置model_type="local"并提供API地址
-
-
第三方API使用:
-
配置相应的API密钥
-
在
ModelConfig中选择对应的model_type
-
-
通过Swagger UI测试:
-
启动应用后访问
http://localhost:8000/docs -
在
/generate-sql端点填写:-
数据库连接信息
-
自然语言查询
-
模型配置
-
-
执行测试
-
三、测试案例
3.1调用大模型连接数据库返回信息
数据库配置
#!/usr/bin/env python
# -*- coding: utf-8 -*-
# @Time : 2020/6/9 14:47
# @Author : CoderCharm
# @File : development_config.py
# @Software: PyCharm
# @Desc :
"""
开发环境配置
"""
from typing import Optional
from pydantic_settings import BaseSettings
class Config(BaseSettings):
# 文档地址
DOCS_URL: str = "/api/v1/docs"
# # 文档关联请求数据接口
OPENAPI_URL: str = "/api/v1/openapi.json"
# 禁用 redoc 文档
REDOC_URL: Optional[str] = None
ACCESS_TOKEN_EXPIRE_MINUTES: int = 60 * 24
JWT_ALGORITHM: str = "HS256"
SECRET_KEY: str = 'koelndom'
# 配置你的Mysql环境
MYSQL_USERNAME: str = 'root'
MYSQL_PASSWORD: str = "root"
MYSQL_HOST: str = "localhost"
# MYSQL_HOST: Union[AnyHttpUrl, IPvAnyAddress] = "119.3.41.115"
MYSQL_DATABASE: str = 'fly'
# Mysql地址
SQLALCHEMY_DATABASE_URI: str = f"mysql+pymysql://{MYSQL_USERNAME}:{MYSQL_PASSWORD}@" \
f"{MYSQL_HOST}/{MYSQL_DATABASE}?charset=utf8"
config = Config()
连接数据库
# !/usr/bin/env python3
# -*- encoding : utf-8 -*-
# @Filename : session.py
# @Software : VSCode
# @Datetime : 2021/11/04 15:49:56
# @Author : leo liu
# @Version : 1.0
# @Description :
from typing import Generator
from sqlalchemy import create_engine
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker
from core.settings import config
engine = create_engine(
config.SQLALCHEMY_DATABASE_URI
)
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
Base = declarative_base()
def get_db() -> Generator:
try:
db = SessionLocal()
yield db
finally:
db.close()
创建API端点
# !/usr/bin/env python3
# -*- encoding : utf-8 -*-
# @Filename : views.py
# @Software : VSCode
# @Datetime : 2021/11/03 17:24:24
# @Author : leo liu
# @Version : 1.0
# @Description :
import json
from typing import Any
from fastapi import APIRouter, Depends, Header
from sqlalchemy.orm.session import Session
import pandas as pd
from extensions import logger
from utils import response_code
from db.session import get_db
from .schemas import ollama_schema
from .crud.ollama import crud_ollama
router = APIRouter()
@router.post("/auth/ollamaRepost", summary="调用大模型返回结果信息")
async def ollama_repost(
*,
db: Session = Depends(get_db),
param: ollama_schema.ollamaBase,
model: ollama_schema.ModelConfig,
) -> Any:
"""
调用大模型返回结果
"""
logger.info(f"查询的问题->:{param.param}")
result = crud_ollama.getresult(db,model,param = param)
logger.info(f"问题解析出的关键字:{result['promptResult']}")
df = pd.DataFrame(result['sqResult'])
print(f"学生信息:\n {df.to_markdown(index=False)}")
logger.info(f"AI分析的结果:\n{result['analysis']}")
return response_code.resp_200(data=result['analysis'], message="success")
定义数据模型
#!/usr/bin/env python
# -*- coding: utf-8 -*-
# @Time : 2020/7/7 16:23
# @Author : CoderCharm
# @File : user_schema.py
# @Software: PyCharm
# @Desc :
"""
"""
from typing import Optional
from pydantic import BaseModel, EmailStr, AnyHttpUrl
class ollamaBase(BaseModel):
param: str
class ModelConfig(BaseModel):
"""大模型配置"""
api_base: str # API基础地址
model_name: str # 模型名称
temperature: float
max_tokens: int
api_key: Optional[str] = None # 可选API密钥
api_version: Optional[str] = None # 某些API需要的版本号
大模型调用的函数
# !/usr/bin/env python3
# -*- encoding : utf-8 -*-
# @Filename : user.py
# @Software : VSCode
# @Datetime : 2021/11/04 21:25:44
# @Author : leo liu
# @Version : 1.0
# @Description :
from pydantic.types import conint
from sqlalchemy import func, desc
from sqlalchemy.orm import Session
from typing import Dict, List
from openai import OpenAI
from sqlalchemy import text
import re
import json
from ..schemas import ollama_schema
from ...etl.test import param
from core.settings import config
import logging
logger = logging.getLogger(__name__)
class CRUDUserAddress():
@staticmethod
def getPromptResult(db: Session, query: ollama_schema.ollamaBase,model:ollama_schema.ModelConfig) -> str:
"""
提取关键字
"""
prompt = f"""
输入文本
{query.param}
提取规则
1. 学号:连续数字,长度1-20位
2. 姓名:2-4个中文字符
返回结果
1. 直接输出JSON对象
2. 不要包含```json等代码块标记
3. 不要有任何额外解释
{{
"code": ["学号",...],
"name":["姓名",...]
}}
"""
return crud_ollama.call_openai_api(model,prompt)
@staticmethod
def get_table_metadata(db: Session) -> Dict:
"""
获取数据库表结构元数据(安全优化版)
参数:
db: SQLAlchemy Session对象
异常:
可能抛出数据库相关异常
"""
metadata = {}
try:
# 1. 获取所有基表
with db.begin() as transaction:
# 使用参数化查询防止SQL注入
tables_result = db.execute(
text("SHOW FULL TABLES WHERE Table_type = 'BASE TABLE'")
).fetchall()
# 动态获取数据库名
db_name = tables_result[0][1] if tables_result else ""
tables = [row[0] for row in tables_result]
logger.info(f"数据库中的表:{tables}")
# 获取每张表的列信息
for table in tables:
columns_result = db.execute(
text(f"SHOW COLUMNS FROM `{table}`") # 直接字符串插值
).fetchall()
columns = [col[0] for col in columns_result]
metadata[table] = {"columns": columns}
logger.info(f"成功获取 {len(metadata)} 张表的元数据")
return metadata
except Exception as e:
logger.error(f"获取元数据失败: {str(e)}")
# 根据业务需求决定是否回滚
if 'transaction' in locals():
transaction.rollback()
raise # 重新抛出异常供上层处理
@staticmethod
def get_sql_result(db:Session,model:ollama_schema.ModelConfig,promptResult,metadata):
"""
获取sql执行结果
"""
print("正在为您生成sql语句并生成结果中---")
# 从提取出的信息中查询想要的结果
code = promptResult["code"]
name = promptResult["name"]
# 关键字查询表和sql
prompt = f"""
已知需要查询的条件:
学生id:{code}
学生姓名:{name}
已知表结构{metadata}
任务要求
1. 分析表结构确定相关的表
2. 找出表之间的关联字段
3. 编写能够联合查询学生基本信息,成绩以及获奖信息的SQL语句,格式规范
4. 直接输出SQL语句内容,不添加任何引号或代码块标记
5. 使用标准SQL语法,保持合理的缩进和换行
6. 返回的字段需设置中文别名并去重
7. 最终结果严格以JSON格式输出,仅包含必要内容:
{{
"related_tables": ["表1", "表2", ...],
"sql": "SELECT ..."
}}
其中JSON结构中键名和字符串值使用双引号,整体不包裹任何额外引号或标记,无多余注释
"""
# 第三步:调用Ollama生成SQL
try:
port = crud_ollama.call_openai_api(model,prompt)
# 获取json格式输出
cleaned_content = re.sub(r'^```json\n|\n```$', '', port.strip())
response_data = json.loads(cleaned_content)
# 获取ollama返回的JSON信息
print(f"AI分析得出的sql结果:\n {port}")
print(f"提取AI获取可能相关的表:\n{response_data.get('related_tables', [])}")
print(f"提取AI获取的sql结果:\n{response_data.get('sql', [])}")
result= crud_ollama.detailSql(db, ''.join(response_data.get('sql', [])).rstrip(';'))
return result
except Exception as e:
return {"error": str(e)}
@staticmethod
def call_openai_api(request: ollama_schema.ModelConfig,prompt) -> str:
"""
调用OpenAI API
参数:
request: 包含prompt和模型配置的请求对象
返回:
模型生成的文本
"""
import openai
try:
client = OpenAI(
base_url=request.api_base, # Ollama的OpenAI兼容端点
api_key="ollama" # 任意非空字符串
)
response = client.chat.completions.create(
model=request.model_name,
messages=[{'role': 'user', 'content': prompt}],
temperature=request.temperature,
max_tokens=request.max_tokens
)
return response.choices[0].message.content
except Exception as e:
print(e)
@staticmethod
def detailSql(db: Session, sql):
try:
# 1. 获取所有基表
with db.begin() as transaction:
# 使用参数化查询防止SQL注入
result = db.execute(
text(sql)
).fetchall()
return result
except Exception as e:
logger.error(f"获取元数据失败: {str(e)}")
# 根据业务需求决定是否回滚
if 'transaction' in locals():
transaction.rollback()
raise # 重新抛出异常供上层处理
@staticmethod
def getresult(db:Session,model:ollama_schema.ModelConfig,param : ollama_schema.ollamaBase):
promptResult = crud_ollama.getPromptResult(db, param, model)
metadata = crud_ollama.get_table_metadata(db)
sqResult = crud_ollama.get_sql_result(db, model, json.loads(promptResult), metadata)
prompt = f"""
针对{sqResult}
总结该学生的学术表现、获奖情况和综合能力,给出200字左右的评价。
"""
analysis = crud_ollama.call_openai_api(model,prompt)
return {
"promptResult":promptResult,
"sqResult":sqResult,
"analysis": analysis
}
crud_ollama = CRUDUserAddress()
本地测试结果
目录结构



3.2、文件上传解析文件内容返回内容需要的结果
准备工作
文件333.txt

创建api端点和调用函数
# !/usr/bin/env python3
# -*- encoding : utf-8 -*-
# @Filename : views.py
# @Software : VSCode
# @Datetime : 2021/11/03 17:24:24
# @Author : leo liu
# @Version : 1.0
# @Description :
import json
from typing import Any
from fastapi import APIRouter, Depends, Header
from sqlalchemy.orm.session import Session
import pandas as pd
from extensions import logger
from utils import response_code
from db.session import get_db
from .schemas import ollama_schema
from .crud.ollama import crud_ollama
from .crud.file import crud_file
from fastapi import FastAPI, UploadFile, File, HTTPException,Form
from typing import Optional
router = APIRouter() #路由分组
@router.post("/auth/fileOllama",summary="解析上传的文件返回结果")
async def ollama_repost(
*,
file: UploadFile = File(..., description="上传的文件"),
config: str = Form(..., description="模型配置(JSON字符串)")
) -> Any:
"""
处理上传的文件并调用大模型
- **file**: 要上传的文件
- **config**: 可选的大模型参数(JSON格式字符串)
"""
try:
# 1. 验证文件
if not file.filename:
raise HTTPException(status_code=400, detail="未提供文件名")
logger.info(f"开始处理文件: {file.filename} (类型: {file.content_type})")
# 2. 解析文件内容
file_content = crud_file.parse_file(file)
logger.info(f"文件解析成功,内容长度: {len(file_content)}字符")
# 3. 调用大模型处理内容
logger.info("调用大模型处理内容...")
model = ollama_schema.ModelConfig.parse_raw(config)
model_response = crud_file.getresult(model,file_content)
logger.info("大模型处理完成")
return response_code.resp_200(data=model_response['analysis'], message="success")
except Exception as e:
logger.info(f"{e}")
# !/usr/bin/env python3
# -*- encoding : utf-8 -*-
# @Filename : pfliu.py
# @Software : VSCode
# @Datetime : 2021/11/04 21:25:44
# @Author : leo liu
# @Version : 1.0
# @Description :
from sqlalchemy.orm import Session
from fastapi import FastAPI, UploadFile, File, HTTPException
from openai import OpenAI
from sqlalchemy import text
import re
import os
import docx2txt
import tempfile
import json
from ..schemas import ollama_schema
from .ollama import crud_ollama
import logging
logger = logging.getLogger(__name__)
class CRUDFile():
@staticmethod
def read_docx_with_docx2txt(file_path):
"""使用docx2txt库解析Word文档"""
try:
return docx2txt.process(file_path)
except Exception as e:
raise ValueError(f"无法解析Word文档: {str(e)}")
@staticmethod
def parse_file(file: UploadFile) -> str:
"""根据文件类型解析文件内容"""
content_type = file.content_type
filename = file.filename
# 创建临时文件
with tempfile.NamedTemporaryFile(delete=False) as temp_file:
temp_file.write(file.file.read())
temp_file_path = temp_file.name
try:
# 根据文件类型选择解析方式
if content_type == "text/plain" or filename.endswith('.txt'):
with open(temp_file_path, 'r', encoding='utf-8') as f:
content = f.read()
elif content_type == "application/json" or filename.endswith('.json'):
import json
with open(temp_file_path, 'r', encoding='utf-8') as f:
json_data = json.load(f)
content = str(json_data)
elif content_type in ["application/vnd.openxmlformats-officedocument.wordprocessingml.document",
"application/msword"] or filename.endswith(('.docx', '.doc')):
doc = crud_file.read_docx_with_docx2txt(temp_file_path)
content = "\n".join([para.text for para in doc.paragraphs])
elif content_type == "application/pdf" or filename.endswith('.pdf'):
import PyPDF2
with open(temp_file_path, 'rb') as f:
reader = PyPDF2.PdfReader(f)
content = "\n".join([page.extract_text() for page in reader.pages])
else:
# 尝试作为文本文件读取
try:
with open(temp_file_path, 'r', encoding='utf-8') as f:
content = f.read()
except:
raise HTTPException(status_code=400, detail="不支持的文件类型")
return content
finally:
# 清理临时文件
try:
os.unlink(temp_file_path)
except:
pass
@staticmethod
def getresult(model:ollama_schema.ModelConfig,file_content):
analysis = crud_ollama.call_openai_api(model,file_content)
return {
"analysis": analysis
}
crud_file = CRUDFile()
swagger调用界面


生成word
# 创建一个 Word 文档
doc = Document()
# 1. 添加标题
doc.add_heading("FastAPI 生成的 Word 文档", level=1)
# 2. 添加正文内容
doc.add_paragraph("这是由 FastAPI 自动生成的 Word 文档内容。")
# 3. 添加"内容"到文档
doc.add_paragraph(model_response['analysis']) # 关键修改:直接添加到文档
# 4. 保存到内存(BytesIO)
file_stream = io.BytesIO() # 创建空BytesIO
doc.save(file_stream) # 将文档写入BytesIO
file_stream.seek(0)
# 返回 Word 文件
return Response(
content=file_stream.getvalue(),
media_type="application/vnd.openxmlformats-officedocument.wordprocessingml.document",
headers={"Content-Disposition": "attachment; filename=output.docx"}
)
四、FastAPI 常用包详解
1. 核心依赖包
fastapi
-
功能:FastAPI 框架本身
-
常用导入:
from fastapi import FastAPI, APIRouter, Request, Response, status, Depends, HTTPException -
关键组件:
-
FastAPI(): 应用实例 -
APIRouter(): 路由分组 -
Depends(): 依赖注入系统
-
uvicorn
-
功能:ASGI 服务器,用于运行 FastAPI 应用
-
使用方式:
uvicorn.run("main:app", host="0.0.0.0", port=8000, reload=True) -
参数:
-
reload: 开发时自动重载 -
workers: 工作进程数
-
2. 请求处理相关包
python-multipart
-
功能:处理表单数据和文件上传
-
使用场景:
from fastapi import File, UploadFile @app.post("/upload/") async def upload_file(file: UploadFile = File(...)): return {"filename": file.filename}
pydantic
-
功能:数据验证和设置管理
-
核心用途:
from pydantic import BaseModel class Item(BaseModel): name: str price: float is_offer: bool = None
3. 数据库集成包
SQLAlchemy
-
功能:ORM 工具
-
常用组合:
-
sqlalchemy: 核心 ORM -
databases: 异步数据库支持
-
-
示例配置:
from sqlalchemy import create_engine from sqlalchemy.ext.declarative import declarative_base SQLALCHEMY_DATABASE_URL = "sqlite:///./sql_app.db" engine = create_engine(SQLALCHEMY_DATABASE_URL) Base = declarative_base()
asyncpg/aiomysql
-
功能:异步 PostgreSQL/MySQL 驱动
-
使用场景:
import asyncpg async def get_db(): return await asyncpg.connect(DATABASE_URL)
4. 认证与安全包
passlib
-
功能:密码哈希
-
常用算法:
from passlib.context import CryptContext pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto") def verify_password(plain_password, hashed_password): return pwd_context.verify(plain_password, hashed_password)
python-jose
-
功能:JWT 实现
-
典型用法:
from jose import JWTError, jwt def create_access_token(data: dict): return jwt.encode(data, SECRET_KEY, algorithm=ALGORITHM)
5. 实用工具包
aiofiles
-
功能:异步文件操作
-
使用示例:
import aiofiles async def save_upload_file(upload_file: UploadFile, destination: Path): async with aiofiles.open(destination, "wb") as buffer: while content := await upload_file.read(1024): await buffer.write(content)
python-multipart
-
功能:处理 multipart 表单数据
-
必要场景:文件上传时自动安装
6. 测试相关包
httpx
-
功能:异步 HTTP 客户端,用于测试
-
测试示例:
from httpx import AsyncClient async def test_create_item(): async with AsyncClient(app=app, base_url="http://test") as ac: response = await ac.post("/items/", json={"name": "Test"}) assert response.status_code == 200
pytest
-
功能:测试框架
-
常用插件:
-
pytest-asyncio: 异步测试支持 -
pytest-cov: 覆盖率测试
-
7. 部署相关包
gunicorn
-
功能:生产环境 WSGI 服务器
-
与 uvicorn 配合使用:
gunicorn -k uvicorn.workers.UvicornWorker -w 4 -b :8000 main:app
python-dotenv
-
功能:环境变量管理
-
使用方式:
from dotenv import load_dotenv load_dotenv() DATABASE_URL = os.getenv("DATABASE_URL")
8. 高级功能包
fastapi-cache2
-
功能:API 响应缓存
-
示例:
from fastapi_cache import FastAPICache from fastapi_cache.backends.redis import RedisBackend FastAPICache.init(RedisBackend(redis), prefix="fastapi-cache")
celery
-
功能:分布式任务队列
-
典型用途:处理耗时任务
from celery import Celery celery = Celery(__name__, broker="redis://localhost:6379/0") @celery.task def process_video(file_path): # 耗时视频处理 pass
更多推荐


所有评论(0)