mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: Hub 连接清理顺序修复、OAuth 类型常量提取、disassociate 逻辑优化
- Hub proxy/worker 连接关闭时先 unregister 再延迟 abort writer,确保缓冲消息排空 - 提取 OAUTH_AUTH_TYPES 常量,替代各处硬编码的 OAuth 类型列表 - auto-disassociate 跳过 OAuth Key,避免其动态 allowed_models 干扰判定 - 删除 Key 时传入 skip_disassociate=True,跳过不必要的解关联检查 - ModelMapper 缓存命中时将 ORM 实例脱离 Session,修复 DetachedInstanceError
This commit is contained in:
2
aether-hub/Cargo.lock
generated
2
aether-hub/Cargo.lock
generated
@@ -10,7 +10,7 @@ checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "aether-hub"
|
name = "aether-hub"
|
||||||
version = "0.1.0"
|
version = "0.1.3"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"axum",
|
"axum",
|
||||||
"clap",
|
"clap",
|
||||||
|
|||||||
@@ -79,11 +79,16 @@ pub async fn handle_proxy_connection(
|
|||||||
.await;
|
.await;
|
||||||
});
|
});
|
||||||
|
|
||||||
// Wait for reader to end, then cleanup writer/ping and unregister from hub.
|
// Wait for reader to end, then cleanup.
|
||||||
let _ = reader.await;
|
let _ = reader.await;
|
||||||
ping_task.abort();
|
ping_task.abort();
|
||||||
writer.abort();
|
|
||||||
hub.unregister_proxy(conn_id, &node_id);
|
hub.unregister_proxy(conn_id, &node_id);
|
||||||
|
// conn still holds an Arc<ProxyConn> with a channel sender clone.
|
||||||
|
// Drop it so the writer can drain and exit.
|
||||||
|
drop(conn);
|
||||||
|
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||||
|
writer.abort();
|
||||||
|
let _ = writer.await;
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn run_proxy_reader(
|
async fn run_proxy_reader(
|
||||||
|
|||||||
@@ -65,11 +65,17 @@ pub async fn handle_worker_connection(
|
|||||||
.await;
|
.await;
|
||||||
});
|
});
|
||||||
|
|
||||||
// Wait for reader to end, then cleanup writer/ping and unregister from hub.
|
// Wait for reader to end, then cleanup.
|
||||||
let _ = reader.await;
|
let _ = reader.await;
|
||||||
ping_task.abort();
|
ping_task.abort();
|
||||||
writer.abort();
|
// Unregister first so all channel senders are dropped (reader_tx dropped
|
||||||
|
// when reader completes, ping_tx dropped by abort, conn.tx dropped when
|
||||||
|
// the last Arc<WorkerConn> is removed from hub). This lets the writer
|
||||||
|
// drain buffered messages (e.g. GOAWAY) before we force-abort it.
|
||||||
hub.unregister_worker(conn_id);
|
hub.unregister_worker(conn_id);
|
||||||
|
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||||
|
writer.abort();
|
||||||
|
let _ = writer.await;
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn run_worker_reader(
|
async fn run_worker_reader(
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ from src.database import create_session
|
|||||||
from src.database.database import get_db
|
from src.database.database import get_db
|
||||||
from src.models.database import Provider, ProviderAPIKey, User
|
from src.models.database import Provider, ProviderAPIKey, User
|
||||||
from src.services.provider.pool.config import parse_pool_config
|
from src.services.provider.pool.config import parse_pool_config
|
||||||
|
from src.services.provider_keys.auth_type import OAUTH_AUTH_TYPES
|
||||||
from src.utils.auth_utils import require_admin
|
from src.utils.auth_utils import require_admin
|
||||||
|
|
||||||
router = APIRouter(prefix="/api/admin/provider-oauth", tags=["Provider OAuth"])
|
router = APIRouter(prefix="/api/admin/provider-oauth", tags=["Provider OAuth"])
|
||||||
@@ -550,7 +551,7 @@ def _check_duplicate_oauth_account(
|
|||||||
# 查询该 Provider 下所有 OAuth 类型的 Keys
|
# 查询该 Provider 下所有 OAuth 类型的 Keys
|
||||||
query = db.query(ProviderAPIKey).filter(
|
query = db.query(ProviderAPIKey).filter(
|
||||||
ProviderAPIKey.provider_id == provider_id,
|
ProviderAPIKey.provider_id == provider_id,
|
||||||
ProviderAPIKey.auth_type.in_(["oauth", "kiro"]), # kiro 也是 OAuth 类型
|
ProviderAPIKey.auth_type.in_(OAUTH_AUTH_TYPES),
|
||||||
)
|
)
|
||||||
if exclude_key_id:
|
if exclude_key_id:
|
||||||
query = query.filter(ProviderAPIKey.id != exclude_key_id)
|
query = query.filter(ProviderAPIKey.id != exclude_key_id)
|
||||||
|
|||||||
@@ -498,14 +498,20 @@ class GlobalModelService:
|
|||||||
logger.warning(f"Provider {provider_id} not found for auto-disassociation")
|
logger.warning(f"Provider {provider_id} not found for auto-disassociation")
|
||||||
return results
|
return results
|
||||||
|
|
||||||
# 1. 先快速检查是否存在“允许所有模型”的活跃 Key。
|
# 1. 先快速检查是否存在"允许所有模型"的活跃 Key。
|
||||||
# 这种情况下无需解除任何关联,避免继续扫描整张 key 表。
|
# 这种情况下无需解除任何关联,避免继续扫描整张 key 表。
|
||||||
|
# 注意:跳过 OAuth Key,OAuth Key 的 allowed_models 由上游动态获取,数量庞大,
|
||||||
|
# 不应参与 disassociate 判定。
|
||||||
|
from src.services.provider_keys.auth_type import OAUTH_AUTH_TYPES
|
||||||
|
|
||||||
|
non_oauth_filter = ProviderAPIKey.auth_type.notin_(OAUTH_AUTH_TYPES)
|
||||||
has_unlimited_key = (
|
has_unlimited_key = (
|
||||||
db.query(ProviderAPIKey.id)
|
db.query(ProviderAPIKey.id)
|
||||||
.filter(
|
.filter(
|
||||||
ProviderAPIKey.provider_id == provider_id,
|
ProviderAPIKey.provider_id == provider_id,
|
||||||
ProviderAPIKey.is_active == True,
|
ProviderAPIKey.is_active == True,
|
||||||
ProviderAPIKey.allowed_models.is_(None),
|
ProviderAPIKey.allowed_models.is_(None),
|
||||||
|
non_oauth_filter,
|
||||||
)
|
)
|
||||||
.limit(1)
|
.limit(1)
|
||||||
.first()
|
.first()
|
||||||
@@ -520,6 +526,7 @@ class GlobalModelService:
|
|||||||
.filter(
|
.filter(
|
||||||
ProviderAPIKey.provider_id == provider_id,
|
ProviderAPIKey.provider_id == provider_id,
|
||||||
ProviderAPIKey.is_active == True,
|
ProviderAPIKey.is_active == True,
|
||||||
|
non_oauth_filter,
|
||||||
)
|
)
|
||||||
.all()
|
.all()
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -110,6 +110,13 @@ class ModelMapperMiddleware:
|
|||||||
)
|
)
|
||||||
|
|
||||||
if model:
|
if model:
|
||||||
|
# 将 ORM Model 转为无 Session 绑定的实例,避免跨请求缓存导致 DetachedInstanceError
|
||||||
|
from sqlalchemy.orm.session import object_session
|
||||||
|
|
||||||
|
if object_session(model) is not None:
|
||||||
|
model_dict = ModelCacheService._model_to_dict(model)
|
||||||
|
model = ModelCacheService._dict_to_model(model_dict)
|
||||||
|
|
||||||
# 创建映射对象
|
# 创建映射对象
|
||||||
mapping = type(
|
mapping = type(
|
||||||
"obj",
|
"obj",
|
||||||
|
|||||||
@@ -2,6 +2,10 @@
|
|||||||
Provider Key 认证类型相关规则。
|
Provider Key 认证类型相关规则。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
# 数据库中所有属于 OAuth 的 auth_type 值(含历史别名)。
|
||||||
|
# 新增 OAuth 类型时只需在此追加,SQL 过滤和 Python 判断均引用此常量。
|
||||||
|
OAUTH_AUTH_TYPES: tuple[str, ...] = ("oauth", "kiro")
|
||||||
|
|
||||||
|
|
||||||
def normalize_auth_type(raw: str) -> str:
|
def normalize_auth_type(raw: str) -> str:
|
||||||
"""将数据库中的 auth_type 归一化为逻辑类型。
|
"""将数据库中的 auth_type 归一化为逻辑类型。
|
||||||
|
|||||||
@@ -134,9 +134,6 @@ async def run_delete_key_side_effects(
|
|||||||
deleted_key_allowed_models: list[str] | None,
|
deleted_key_allowed_models: list[str] | None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""执行删除 Key 后的副作用。"""
|
"""执行删除 Key 后的副作用。"""
|
||||||
# 触发缓存失效和自动解除关联检查
|
|
||||||
# 注意:删除后是否需要解除关联,应基于“删除后的活跃 Key 集合”判断。
|
|
||||||
# 不能仅凭被删除 Key 的 allowed_models 是否为 null 来跳过 disassociate。
|
|
||||||
_ = deleted_key_allowed_models
|
_ = deleted_key_allowed_models
|
||||||
if provider_id:
|
if provider_id:
|
||||||
from src.services.model.global_model import on_key_allowed_models_changed
|
from src.services.model.global_model import on_key_allowed_models_changed
|
||||||
@@ -144,6 +141,7 @@ async def run_delete_key_side_effects(
|
|||||||
await on_key_allowed_models_changed(
|
await on_key_allowed_models_changed(
|
||||||
db=db,
|
db=db,
|
||||||
provider_id=provider_id,
|
provider_id=provider_id,
|
||||||
|
skip_disassociate=True,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# 无 provider_id 时仅清除缓存
|
# 无 provider_id 时仅清除缓存
|
||||||
|
|||||||
@@ -202,7 +202,7 @@ def test_clear_oauth_invalid_response_invalidates_caches(
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_run_delete_key_side_effects_not_skip_disassociate(
|
async def test_run_delete_key_side_effects_skip_disassociate(
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
) -> None:
|
) -> None:
|
||||||
captured: dict[str, Any] = {}
|
captured: dict[str, Any] = {}
|
||||||
@@ -225,4 +225,4 @@ async def test_run_delete_key_side_effects_not_skip_disassociate(
|
|||||||
)
|
)
|
||||||
|
|
||||||
assert captured["provider_id"] == "provider-1"
|
assert captured["provider_id"] == "provider-1"
|
||||||
assert "skip_disassociate" not in captured
|
assert captured["skip_disassociate"] is True
|
||||||
|
|||||||
Reference in New Issue
Block a user