feat: 初始化工程

This commit is contained in:
2026-09-06 16:07:00 +08:00
commit 30b5bf2aa1
36 changed files with 1166 additions and 0 deletions

68
.gitignore vendored Normal file
View File

@@ -0,0 +1,68 @@
# Python 字节码文件
__pycache__/
*.py[cod]
*$py.class
# C 扩展
*.so
# 分发/打包
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
*.egg-info/
.installed.cfg
*.egg
# 虚拟环境
venv/
env/
ENV/
.env
.venv
# 测试
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
.hypothesis/
# Django 相关
*.log
local_settings.py
db.sqlite3
db.sqlite3-journal
media/
# PyCharm IDE
.idea/
*.iml
*.iws
*.ipr
# VS Code
.vscode/
*.code-workspace
.history/
# 其他
.DS_Store
logs/
packages/

10
.idea/.gitignore generated vendored Normal file
View File

@@ -0,0 +1,10 @@
# Default ignored files
/shelf/
/workspace.xml
# Ignored default folder with query files
/queries/
# Datasource local storage ignored files
/dataSources/
/dataSources.local.xml
# Editor-based HTTP Client requests
/httpRequests/

10
.idea/family-service.iml generated Normal file
View File

@@ -0,0 +1,10 @@
<?xml version="1.0" encoding="UTF-8"?>
<module type="PYTHON_MODULE" version="4">
<component name="NewModuleRootManager">
<content url="file://$MODULE_DIR$">
<excludeFolder url="file://$MODULE_DIR$/.venv" />
</content>
<orderEntry type="jdk" jdkName="Python 3.12 (family-service)" jdkType="Python SDK" />
<orderEntry type="sourceFolder" forTests="false" />
</component>
</module>

View File

@@ -0,0 +1,6 @@
<component name="InspectionProjectProfileManager">
<settings>
<option name="USE_PROJECT_PROFILE" value="false" />
<version value="1.0" />
</settings>
</component>

8
.idea/modules.xml generated Normal file
View File

@@ -0,0 +1,8 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="ProjectModuleManager">
<modules>
<module fileurl="file://$PROJECT_DIR$/.idea/family-service.iml" filepath="$PROJECT_DIR$/.idea/family-service.iml" />
</modules>
</component>
</project>

6
.idea/vcs.xml generated Normal file
View File

@@ -0,0 +1,6 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="VcsDirectoryMappings">
<mapping directory="$PROJECT_DIR$" vcs="Git" />
</component>
</project>

11
Dockerfile Normal file
View File

@@ -0,0 +1,11 @@
FROM python:3.12-slim
WORKDIR /app
RUN ln -sf /usr/share/zoneinfo/Asia/Shanghai /etc/localtime
RUN echo 'Asia/Shanghai' > /etc/timezone
COPY ./packages /app/packages
COPY requirements.txt /app/
RUN pip install --no-cache-dir --no-index --find-links=/app/packages -r requirements.txt
COPY . /app/
EXPOSE 8000
CMD ["uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000"]
# pip download -r requirements.txt -d ./packages --only-binary=:all: --platform manylinux2014_x86_64 -i https://pypi.tuna.tsinghua.edu.cn/simple

26
config/database.py Normal file
View File

@@ -0,0 +1,26 @@
import os
from dotenv import load_dotenv
from sqlalchemy import create_engine
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker
load_dotenv()
DATABASE_URL = f"mysql+pymysql://root:{os.getenv('DB_PASSWORD')}@{os.getenv('DB_HOST')}:3306/family"
engine = create_engine(url=DATABASE_URL, pool_pre_ping=True, pool_recycle=3600)
# 会话工厂
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
# ORM基类
Base = declarative_base()
def get_db():
# 创建数据库会话实例
db = SessionLocal()
try:
yield db
finally:
db.close()

18
config/rustfs.py Normal file
View File

@@ -0,0 +1,18 @@
import os
import boto3
from botocore.client import Config
from dotenv import load_dotenv
load_dotenv()
access_key = 'TRatSfBovO0W36NUhP2c'
secret_access = '4ZFzdOUjSVeNu0W8RGg5MasyBmpL79rlAHQwb32Y'
s3 = boto3.client('s3',
endpoint_url=f'http://{os.getenv("RUSTFS_HOST")}:{os.getenv("RUSTFS_PORT")}',
aws_access_key_id=access_key,
aws_secret_access_key=secret_access,
config=Config(signature_version='s3v4'),
region_name='cn-east-1'
)

0
id_generator/__init__.py Normal file
View File

38
id_generator/generator.py Normal file
View File

@@ -0,0 +1,38 @@
"""
雪花算法生成器IdGenerator
"""
# !/usr/bin/python
# coding=UTF-8
from . import options
from . import snowflake_m1
class DefaultIdGenerator:
"""
ID生成器
"""
def __init__(self):
self.snowflake = None
def set_id_generator(self, option: options.IdGeneratorOptions):
"""
设置id生成规则信息
"""
if option.base_time < 100000:
raise ValueError("base time error.")
self.snowflake = snowflake_m1.SnowFlakeM1(option)
def next_id(self) -> int:
"""
获取新的UUID
"""
if self.snowflake is None:
raise ValueError("please set id generator at first.")
return self.snowflake.next_id()

134
id_generator/idregister.py Normal file
View File

@@ -0,0 +1,134 @@
"""
worker id generator
"""
# !/usr/bin/python
# coding=UTF-8
from threading import Thread
import time
import logging
import redis
class Register:
"""
redis封装
- host 代表redis ip
- port 代表redis端口
- max_worker_id worker_id的最大值, 默认为100
- password redis的密码, 默认为空
"""
def __init__(self, host, port, max_worker_id=100, password=None):
self.redis_impl = redis.StrictRedis(host=host, port=port, db=0, password=password)
self.loop_count = 0
self.max_loop_count = 10
self.worker_id_expire_time = 15
self.max_worker_id = max_worker_id
self.worker_id = -1
self.is_stop = False
def get_lock(self, key):
"""
获取分布式全局锁,并设置过期时间为30秒
"""
if self.redis_impl.setnx(key, 1):
self.redis_impl.expire(key, 30)
return True
if self.redis_impl.ttl(key) < 0:
self.redis_impl.expire(key, 30)
return False
def stop(self):
"""
退出注册器的线程
"""
self.is_stop = True
def get_worker_id(self):
"""
获取全局唯一worker_id, 会创建一个线程给worker id续期
失败返回-1
"""
self.loop_count = 0
def extern_life(my_id):
while 1:
time.sleep(self.worker_id_expire_time / 3)
# 是否关闭了
if self.is_stop:
return
# 更新生命周期
if self.worker_id != my_id:
break
try:
self.redis_impl.expire(
f"IdGen:WorkerId:Value:{my_id}",
self.worker_id_expire_time)
except Exception as exe:
logging.error(exe)
continue
self.worker_id = self.__get_next_worker_id()
if self.worker_id > -1:
Thread(target=extern_life, args=[self.worker_id]).start()
return self.worker_id
def __get_next_worker_id(self):
"""
获取全局唯一worker id内部实现
"""
cur = self.redis_impl.incrby("IdGen:WorkerId:Index", 1)
def can_reset():
try:
reset_value = self.redis_impl.incr("IdGen:WorkerId:Value:Edit")
return reset_value != 1
except Exception as ept:
logging.error(ept)
return False
def end_reset():
try:
self.redis_impl.set("IdGen:WorkerId:Value:Edit", 0)
except Exception as ept:
logging.error(ept)
def is_available(worker_id: int):
try:
rst = self.redis_impl.get(f"IdGen:WorkerId:Value:{worker_id}")
return rst != "Y"
except Exception as ept:
logging.error(ept)
return False
if cur > self.max_worker_id:
if can_reset():
self.redis_impl.set("IdGen:WorkerId:Index", -1)
end_reset()
self.loop_count += 1
if self.loop_count > self.max_loop_count:
self.loop_count = 0
return -1
time.sleep(0.2 * self.loop_count)
return self.__get_next_worker_id()
time.sleep(0.2)
return self.__get_next_worker_id()
if is_available(cur):
self.redis_impl.setex(
f"IdGen:WorkerId:Value:{cur}",
self.worker_id_expire_time,
"Y"
)
self.loop_count = 0
return cur
return self.__get_next_worker_id()

43
id_generator/options.py Normal file
View File

@@ -0,0 +1,43 @@
"""
生成器IdGenerator配置选项
"""
# !/usr/bin/python
# coding=UTF-8
class IdGeneratorOptions:
"""
ID生成器配置
- worker_id 全局唯一id, 区分不同uuid生成器实例
- worker_id_bit_length 生成的uuid中worker_id占用的位数
- seq_bit_length 生成的uuid中序列号占用的位数
"""
def __init__(self, worker_id=0, worker_id_bit_length=6, seq_bit_length=6):
# 雪花计算方法,1-漂移算法|2-传统算法), 默认1。目前只实现了1。
self.method = 1
# 基础时间ms单位, 不能超过当前系统时间
self.base_time = 1582136402000
# 机器码, 必须由外部设定, 最大值 2^worker_id_bit_length-1
self.worker_id = worker_id
# 机器码位长, 默认值6, 取值范围 [1, 15](要求:序列数位长+机器码位长不超过22
self.worker_id_bit_length = worker_id_bit_length
# 序列数位长, 默认值6, 取值范围 [3, 21](要求:序列数位长+机器码位长不超过22
self.seq_bit_length = seq_bit_length
# 最大序列数(含), 设置范围 [max_seq_number, 2^seq_bit_length-1]
# 默认值0, 表示最大序列数取最大值2^seq_bit_length-1]
self.max_seq_number = 0
# 最小序列数(含), 默认值5, 取值范围 [5, max_seq_number], 每毫秒的前5个序列数对应编号0-4是保留位
# 其中1-4是时间回拨相应预留位, 0是手工新值预留位
self.min_seq_number = 5
# 最大漂移次数(含), 默认2000, 推荐范围500-10000与计算能力有关
self.top_over_cost_count = 2000

20
id_generator/snowflake.py Normal file
View File

@@ -0,0 +1,20 @@
"""
雪花算法生成器接口声明
"""
# !/usr/bin/python
# coding=UTF-8
class SnowFlake():
def __init__(self, options):
self.options = options
def next_id(self) -> int:
"""
获取新的UUID
"""
return 0

View File

@@ -0,0 +1,147 @@
"""
M1生成器
"""
# !/usr/bin/python
# coding=UTF-8
import threading
import time
from .snowflake import SnowFlake
from .options import IdGeneratorOptions
class SnowFlakeM1(SnowFlake):
"""
M1规则ID生成器配置
"""
def __init__(self, options: IdGeneratorOptions):
# 1.base_time
self.base_time = 1582136402000
if options.base_time != 0:
self.base_time = int(options.base_time)
# 2.worker_id_bit_length
self.worker_id_bit_length = 6
if options.worker_id_bit_length != 0:
self.worker_id_bit_length = int(options.worker_id_bit_length)
# 3.worker_id
self.worker_id = options.worker_id
# 4.seq_bit_length
self.seq_bit_length = 6
if options.seq_bit_length != 0:
self.seq_bit_length = int(options.seq_bit_length)
# 5.max_seq_number
self.max_seq_number = int(options.max_seq_number)
if options.max_seq_number <= 0:
self.max_seq_number = (1 << self.seq_bit_length) - 1
# 6.min_seq_number
self.min_seq_number = int(options.min_seq_number)
# 7.top_over_cost_count
self.top_over_cost_count = int(options.top_over_cost_count)
# 8.Others
self.__timestamp_shift = self.worker_id_bit_length + self.seq_bit_length
self.__current_seq_number = self.min_seq_number
self.__last_time_tick: int = 0
self.__turn_back_time_tick: int = 0
self.__turn_back_index: int = 0
self.__is_over_cost = False
self.___over_cost_count_in_one_term: int = 0
self.__id_lock = threading.Lock()
def __next_over_cost_id(self) -> int:
current_time_tick = self.__get_current_time_tick()
if current_time_tick > self.__last_time_tick:
self.__last_time_tick = current_time_tick
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = False
self.___over_cost_count_in_one_term = 0
return self.__calc_id(self.__last_time_tick)
if self.___over_cost_count_in_one_term >= self.top_over_cost_count:
self.__last_time_tick = self.__get_next_time_tick()
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = False
self.___over_cost_count_in_one_term = 0
return self.__calc_id(self.__last_time_tick)
if self.__current_seq_number > self.max_seq_number:
self.__last_time_tick += 1
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = True
self.___over_cost_count_in_one_term += 1
return self.__calc_id(self.__last_time_tick)
return self.__calc_id(self.__last_time_tick)
def __next_normal_id(self) -> int:
current_time_tick = self.__get_current_time_tick()
if current_time_tick < self.__last_time_tick:
if self.__turn_back_time_tick < 1:
self.__turn_back_time_tick = self.__last_time_tick - 1
self.__turn_back_index += 1
# 每毫秒序列数的前5位是预留位, 0用于手工新值, 1-4是时间回拨次序
# 支持4次回拨次序避免回拨重叠导致ID重复, 可无限次回拨(次序循环使用)。
if self.__turn_back_index > 4:
self.__turn_back_index = 1
return self.__calc_turn_back_id(self.__turn_back_time_tick)
# 时间追平时, _TurnBackTimeTick清零
self.__turn_back_time_tick = min(self.__turn_back_time_tick, 0)
if current_time_tick > self.__last_time_tick:
self.__last_time_tick = current_time_tick
self.__current_seq_number = self.min_seq_number
return self.__calc_id(self.__last_time_tick)
if self.__current_seq_number > self.max_seq_number:
self.__last_time_tick += 1
self.__current_seq_number = self.min_seq_number
self.__is_over_cost = True
self.___over_cost_count_in_one_term = 1
return self.__calc_id(self.__last_time_tick)
return self.__calc_id(self.__last_time_tick)
def __calc_id(self, use_time_tick) -> int:
self.__current_seq_number += 1
return (
(use_time_tick << self.__timestamp_shift) +
(self.worker_id << self.seq_bit_length) +
self.__current_seq_number
) % int(1e64)
def __calc_turn_back_id(self, use_time_tick) -> int:
self.__turn_back_time_tick -= 1
return (
(use_time_tick << self.__timestamp_shift) +
(self.worker_id << self.seq_bit_length) +
self.__turn_back_index
) % int(1e64)
def __get_current_time_tick(self) -> int:
return int((time.time_ns() / 1e6) - self.base_time)
def __get_next_time_tick(self) -> int:
temp_time_ticker = self.__get_current_time_tick()
while temp_time_ticker <= self.__last_time_tick:
# 0.001 = 1 mili sec
time.sleep(0.001)
temp_time_ticker = self.__get_current_time_tick()
return temp_time_ticker
def next_id(self) -> int:
with self.__id_lock:
if self.__is_over_cost:
nextid = self.__next_over_cost_id()
else:
nextid = self.__next_normal_id()
return nextid

16
main.py Normal file
View File

@@ -0,0 +1,16 @@
from fastapi import FastAPI
from starlette.middleware.cors import CORSMiddleware
from routers import routers
app = FastAPI(title="Family Service")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
for router in routers:
app.include_router(router)

41
models/base.py Normal file
View File

@@ -0,0 +1,41 @@
from sqlalchemy import Column, BigInteger, DateTime, event
from sqlalchemy.ext.declarative import declared_attr
from config.database import Base
from datetime import datetime
from utils.common import camel_to_snake
from id_generator import options, generator
# https://github.com/yitter/IdGenerator/tree/master/Python
options = options.IdGeneratorOptions(worker_id=23)
idgen = generator.DefaultIdGenerator()
idgen.set_id_generator(options)
# 第二层基类包含ID
class IdBase(Base):
__abstract__ = True
id = Column(BigInteger, primary_key=True, index=True)
@declared_attr
def __tablename__(cls):
# 自动把数据库实体类名驼峰转为数据库表名下划线
return camel_to_snake(cls.__name__)
# 自动填充id
@event.listens_for(IdBase, 'before_insert', propagate=True)
def before_insert_listener(mapper, connection, target):
if target.id is None:
target.id = idgen.next_id()
# 第二层基类包含ID和审计字段
class AuditBase(IdBase):
__abstract__ = True
create_time = Column(DateTime, nullable=True, default=datetime.now)
update_time = Column(DateTime, nullable=True, default=datetime.now, onupdate=datetime.now)

41
models/family.py Normal file
View File

@@ -0,0 +1,41 @@
from sqlalchemy import Column, JSON, String, Text, Integer, BigInteger, ForeignKey
from sqlalchemy.orm import relationship, Mapped
from models.base import AuditBase
class Users(AuditBase):
name = Column(String(45), nullable=False, comment="昵称")
avatar = Column(String(255), nullable=True, comment="头像")
records: Mapped[list["Records"]] = relationship("Records", back_populates="user")
class Records(AuditBase):
user_id = Column(BigInteger, ForeignKey("users.id"), nullable=False, comment="用户id")
content = Column(Text, nullable=False, comment="内容")
image_list = Column(JSON, nullable=True, default=list, comment="图片")
video_list = Column(JSON, nullable=True, default=list, comment="视频")
like_count = Column(Integer, nullable=True, default=0, comment="点赞数")
comment_count = Column(Integer, nullable=True, default=0, comment="评论数")
user: Mapped["Users"] = relationship("Users", back_populates="records")
comments: Mapped[list["Comments"]] = relationship(
"Comments",
back_populates="record",
lazy="selectin",
)
class Likes(AuditBase):
record_id = Column(BigInteger, ForeignKey("records.id"), nullable=False, comment="记录id")
user_id = Column(BigInteger, ForeignKey("users.id"), nullable=False, comment="用户id")
class Comments(AuditBase):
record_id = Column(BigInteger, ForeignKey("records.id"), nullable=False, comment="记录id")
user_id = Column(BigInteger, ForeignKey("users.id"), nullable=False, comment="用户id")
content = Column(String(255), nullable=False, comment="内容")
record: Mapped["Records"] = relationship("Records", back_populates="comments")
user: Mapped["Users"] = relationship("Users") # ← 就加了这一行

12
requirements.txt Normal file
View File

@@ -0,0 +1,12 @@
fastapi~=0.140.0
python-dotenv~=1.2.2
requests~=2.34.2
starlette~=1.3.1
pydantic~=2.13.4
SQLAlchemy~=2.0.51
asyncpg~=0.30.0
uvicorn~=0.23.0
pymysql~=1.2.0
boto3~=1.40.59
botocore~=1.40.59
python-multipart~=0.0.20

13
routers/__init__.py Normal file
View File

@@ -0,0 +1,13 @@
from .user import router as user_router
from .record import router as record_router
from .like import router as like_router
from .comment import router as comment_router
from .file import router as file_router
routers = [
user_router,
record_router,
like_router,
comment_router,
file_router
]

23
routers/comment.py Normal file
View File

@@ -0,0 +1,23 @@
from fastapi import APIRouter, Depends
from sqlalchemy.orm import Session
from config.database import get_db
from schemas.comment import CommentCreate, CommentUpdate
from service import comment_service
router = APIRouter(prefix="/comments", tags=["评论"])
@router.post("", response_model=bool, summary="发表评论")
def create(comment_in: CommentCreate, db: Session = Depends(get_db)):
return comment_service.create_comment(db, comment_in)
@router.put("/{comment_id}", response_model=bool, summary="编辑评论")
def update(comment_id: int, comment_in: CommentUpdate, db: Session = Depends(get_db)):
return comment_service.update_comment(db, comment_id, comment_in)
@router.delete("/{comment_id}", summary="删除评论")
def delete(comment_id: int, current_user_id: int, db: Session = Depends(get_db)):
return comment_service.delete_comment(db, comment_id, current_user_id)

10
routers/file.py Normal file
View File

@@ -0,0 +1,10 @@
from fastapi import APIRouter, Query, UploadFile, File
from service import file_service
router = APIRouter(prefix="/file", tags=["文件"])
@router.post("/upload", summary="上传文件", response_model=str)
async def upload(md5: str = Query(..., description="文件MD5值"),
file: UploadFile = File(..., description="要上传的文件")):
return await file_service.upload_file(md5, file)

18
routers/like.py Normal file
View File

@@ -0,0 +1,18 @@
from fastapi import APIRouter, Depends
from sqlalchemy.orm import Session
from config.database import get_db
from schemas.like import LikeCreate
from service import like_service
router = APIRouter(prefix="/likes", tags=["点赞"])
@router.post("", summary="点赞")
def like(like_in: LikeCreate, db: Session = Depends(get_db)):
return like_service.create_like(db, like_in)
@router.delete("", summary="取消点赞")
def unlike(record_id: int, user_id: int, db: Session = Depends(get_db)):
return like_service.delete_like(db, record_id, user_id)

37
routers/record.py Normal file
View File

@@ -0,0 +1,37 @@
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from typing import List
from config.database import get_db
from schemas.record import RecordCreate, RecordUpdate, RecordResponse
from service import record_service
router = APIRouter(prefix="/records", tags=["记录"])
@router.post("", response_model=bool, summary="发布记录")
def create(record_in: RecordCreate, db: Session = Depends(get_db)):
return record_service.create_record(db, record_in)
@router.get("", response_model=List[RecordResponse], summary="获取所有记录")
def list_all(skip: int = 0, limit: int = 20, db: Session = Depends(get_db)):
return record_service.get_all_records(db, skip, limit)
@router.get("/{record_id}", response_model=RecordResponse, summary="获取记录详情")
def get(record_id: int, db: Session = Depends(get_db)):
record = record_service.get_record_by_id(db, record_id)
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
return record
@router.put("/{record_id}", response_model=bool, summary="更新记录")
def update(record_id: int, record_in: RecordUpdate, db: Session = Depends(get_db)):
return record_service.update_record(db, record_id, record_in)
@router.delete("/{record_id}", summary="删除记录")
def delete(record_id: int, db: Session = Depends(get_db)):
return record_service.delete_record(db, record_id)

31
routers/user.py Normal file
View File

@@ -0,0 +1,31 @@
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from config.database import get_db
from schemas.user import UserCreate, UserUpdate, UserResponse
from service import user_service
router = APIRouter(prefix="/users", tags=["用户"])
@router.post("", response_model=UserResponse, summary="创建用户")
def create(user_in: UserCreate, db: Session = Depends(get_db)):
return user_service.create_user(db, user_in)
@router.get("/{user_id}", response_model=UserResponse, summary="获取用户详情")
def get(user_id: int, db: Session = Depends(get_db)):
user = user_service.get_user_by_id(db, user_id)
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
return user
@router.put("/{user_id}", response_model=UserResponse, summary="更新用户")
def update(user_id: int, user_in: UserUpdate, db: Session = Depends(get_db)):
return user_service.update_user(db, user_id, user_in)
@router.delete("/{user_id}", summary="删除用户")
def delete(user_id: int, db: Session = Depends(get_db)):
return user_service.delete_user(db, user_id)

31
schemas/comment.py Normal file
View File

@@ -0,0 +1,31 @@
from datetime import datetime
from pydantic import BaseModel
from schemas.user import UserResponse
class CommentCreate(BaseModel):
record_id: int
user_id: int
content: str
class CommentUpdate(BaseModel):
content: str
class CommentResponse(BaseModel):
id: int
record_id: int
user: UserResponse
content: str
create_time: datetime
update_time: datetime
model_config = {
"from_attributes": True,
"json_encoders": {
datetime: lambda dt: dt.strftime("%Y-%m-%d %H:%M:%S")
}
}

6
schemas/like.py Normal file
View File

@@ -0,0 +1,6 @@
from pydantic import BaseModel
class LikeCreate(BaseModel):
record_id: int
user_id: int

40
schemas/record.py Normal file
View File

@@ -0,0 +1,40 @@
from datetime import datetime
from pydantic import BaseModel
from typing import Optional, List
from schemas.comment import CommentResponse
from schemas.user import UserResponse
class RecordCreate(BaseModel):
user_id: int
content: str
image_list: List[str] = []
video_list: List[str] = []
class RecordUpdate(BaseModel):
content: Optional[str] = None
image_list: Optional[List[str]] = None
video_list: Optional[List[str]] = None
class RecordResponse(BaseModel):
id: int
user: UserResponse
content: str
image_list: List[str]
video_list: List[str]
like_count: int
comment_count: int
comments: List[CommentResponse] = []
create_time: datetime
update_time: datetime
model_config = {
"from_attributes": True,
"json_encoders": {
datetime: lambda dt: dt.strftime("%Y-%m-%d %H:%M:%S")
}
}

29
schemas/user.py Normal file
View File

@@ -0,0 +1,29 @@
from datetime import datetime
from pydantic import BaseModel
from typing import Optional
class UserCreate(BaseModel):
name: str
avatar: Optional[str] = None
class UserUpdate(BaseModel):
name: Optional[str] = None
avatar: Optional[str] = None
class UserResponse(BaseModel):
id: int
name: str
avatar: Optional[str] = None
create_time: datetime
update_time: datetime
model_config = {
"from_attributes": True,
"json_encoders": {
datetime: lambda dt: dt.strftime("%Y-%m-%d %H:%M:%S")
}
}

View File

@@ -0,0 +1,53 @@
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.family import Comments
from schemas.comment import CommentCreate, CommentUpdate
from service.record_service import get_record_by_id
def get_comment_by_id(db: Session, comment_id: int) -> Comments | None:
result = db.execute(select(Comments).where(Comments.id == comment_id))
return result.scalar_one_or_none()
def create_comment(db: Session, comment_dto: CommentCreate) -> bool:
comment = Comments(**comment_dto.model_dump())
db.add(comment)
record = get_record_by_id(db, comment_dto.record_id)
if record:
record.comment_count = (record.comment_count or 0) + 1
db.commit()
db.refresh(comment)
return True
def update_comment(db: Session, comment_id: int, comment_dto: CommentUpdate) -> bool:
comment = get_comment_by_id(db, comment_id)
if not comment:
return False
comment.content = comment_dto.content
db.commit()
db.refresh(comment)
return True
def delete_comment(db: Session, comment_id: int, current_user_id: int) -> bool:
comment = get_comment_by_id(db, comment_id)
if not comment:
return False
if comment.user_id != current_user_id:
return False
record = get_record_by_id(db, comment.record_id)
if record and record.comment_count and record.comment_count > 0:
record.comment_count -= 1
db.delete(comment)
db.commit()
return True

48
service/file_service.py Normal file
View File

@@ -0,0 +1,48 @@
from fastapi import UploadFile, File, HTTPException
from config.rustfs import s3
ALLOWED_IMAGE_TYPES = [
# 图片
"image/jpeg",
"image/png",
"image/gif",
"image/webp",
"image/svg+xml",
# 视频
"video/mp4",
"video/mpeg",
"video/quicktime",
"video/x-msvideo",
"video/webm",
"video/x-matroska",
"video/ogg",
]
BUCKET = 'family'
NGINX_PROXY = 'rustfs'
async def upload_file(md5: str, file: UploadFile = File(...)) -> str:
# 校验文件类型
if file.content_type not in ALLOWED_IMAGE_TYPES:
raise HTTPException(
status_code=400,
detail="只允许上传图片或视频文件 (JPEG, PNG, GIF, WEBP, SVG, MP4, MOV, AVI, WEBM, MKV, OGV)"
)
file_ext = file.filename.split('.')[-1]
unique_filename = f"{md5}.{file_ext}"
file_content = await file.read()
# 上传到S3
s3.put_object(
Bucket=BUCKET,
Key=unique_filename,
Body=file_content,
ContentType=file.content_type
)
# 返回文件url
return f"{NGINX_PROXY}/{BUCKET}/{unique_filename}"

46
service/like_service.py Normal file
View File

@@ -0,0 +1,46 @@
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.family import Likes
from schemas.like import LikeCreate
from service.record_service import get_record_by_id
def get_like(db: Session, record_id: int, user_id: int) -> Likes | None:
result = db.execute(
select(Likes).where(Likes.record_id == record_id, Likes.user_id == user_id)
)
return result.scalar_one_or_none()
def create_like(db: Session, like_dto: LikeCreate) -> bool:
exist = get_like(db, like_dto.record_id, like_dto.user_id)
if exist:
return False
like = Likes(**like_dto.model_dump())
db.add(like)
record = get_record_by_id(db, like_dto.record_id)
if record:
record.like_count = (record.like_count or 0) + 1
db.commit()
return True
def delete_like(db: Session, record_id: int, user_id: int) -> bool:
like = get_like(db, record_id, user_id)
if not like:
return False
db.delete(like)
record = get_record_by_id(db, record_id)
if record and record.like_count and record.like_count > 0:
record.like_count -= 1
db.commit()
return True

59
service/record_service.py Normal file
View File

@@ -0,0 +1,59 @@
from sqlalchemy import select, desc
from sqlalchemy.orm import Session, selectinload
from models.family import Records, Comments
from schemas.record import RecordCreate, RecordUpdate, RecordResponse
def get_all_records(db: Session, skip: int = 0, limit: int = 20) -> list[RecordResponse]:
result = db.execute(
select(Records)
.options(
selectinload(Records.user),
selectinload(Records.comments).selectinload(Comments.user),
)
.order_by(desc(Records.create_time))
.offset(skip)
.limit(limit)
)
records = result.scalars().all()
return [RecordResponse.model_validate(r) for r in records]
def create_record(db: Session, record_dto: RecordCreate) -> bool:
data = record_dto.model_dump()
record = Records(**data)
db.add(record)
db.commit()
db.refresh(record)
return True
def get_record_by_id(db: Session, record_id: int) -> Records | None:
result = db.execute(select(Records).where(Records.id == record_id))
return result.scalar_one_or_none()
def update_record(db: Session, record_id: int, record_dto: RecordUpdate) -> bool:
record = get_record_by_id(db, record_id)
if not record:
return False
for field, value in record_dto.model_dump(exclude_unset=True).items():
setattr(record, field, value)
db.commit()
db.refresh(record)
return True
def delete_record(db: Session, record_id: int) -> bool:
record = get_record_by_id(db, record_id)
if not record:
return False
db.delete(record)
db.commit()
return True

43
service/user_service.py Normal file
View File

@@ -0,0 +1,43 @@
from sqlalchemy import select
from sqlalchemy.orm import Session
from models.family import Users
from schemas.user import UserCreate, UserUpdate
def get_user_by_id(db: Session, user_id: int) -> Users | None:
result = db.execute(select(Users).where(Users.id == user_id))
return result.scalar_one_or_none()
def create_user(db: Session, user: UserCreate) -> Users:
user = Users(**user.model_dump())
db.add(user)
db.commit()
db.refresh(user)
return user
def update_user(db: Session, user_id: int, obj_in: UserUpdate) -> Users:
user = get_user_by_id(db, user_id)
if not user:
return False
for field, value in obj_in.model_dump(exclude_unset=True).items():
setattr(user, field, value)
db.commit()
db.refresh(user)
return user
def delete_user(db: Session, user_id: int) -> bool:
user = get_user_by_id(db, user_id)
if not user:
return False
db.delete(user)
db.commit()
return True

11
test_main.http Normal file
View File

@@ -0,0 +1,11 @@
# Test your FastAPI endpoints
GET http://127.0.0.1:8000/
Accept: application/json
###
GET http://127.0.0.1:8000/hello/User
Accept: application/json
###

13
utils/common.py Normal file
View File

@@ -0,0 +1,13 @@
import re
def camel_to_snake(name: str) -> str:
"""将驼峰命名转换为蛇形命名CamelCase → snake_case"""
name = re.sub('(.)([A-Z][a-z]+)', r'\1_\2', name)
return re.sub('([a-z0-9])([A-Z])', r'\1_\2', name).lower()
def snake_to_camel(name: str) -> str:
"""将蛇形命名转换为驼峰命名snake_case → CamelCase"""
components = name.split('_')
return ''.join(x.title() for x in components)