Skip to content

Commit 7611878

Browse files
committed
[ty] diagnostic on overridden __setattr__ and __delattr__ in frozen dataclasses
astral-sh/ty#111
1 parent 2d0681d commit 7611878

7 files changed

Lines changed: 269 additions & 142 deletions

File tree

crates/ty/docs/rules.md

Lines changed: 106 additions & 74 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

crates/ty_python_semantic/resources/mdtest/dataclasses/dataclasses.md

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -443,7 +443,7 @@ frozen_instance = MyFrozenClass(1)
443443
frozen_instance.x = 2 # error: [invalid-assignment]
444444
```
445445

446-
If `__setattr__()` or `__delattr__()` is defined in the class, we should emit a diagnostic.
446+
If `__setattr__()` or `__delattr__()` is defined in the class, a diagnostic is emitted.
447447

448448
```py
449449
from dataclasses import dataclass
@@ -452,10 +452,10 @@ from dataclasses import dataclass
452452
class MyFrozenClass:
453453
x: int
454454

455-
# TODO: Emit a diagnostic here
455+
# error: [unsound-dataclass-method-override] "Cannot overwrite attribute `__setattr__` in class `MyFrozenClass`"
456456
def __setattr__(self, name: str, value: object) -> None: ...
457457

458-
# TODO: Emit a diagnostic here
458+
# error: [unsound-dataclass-method-override] "Cannot overwrite attribute `__delattr__` in class `MyFrozenClass`"
459459
def __delattr__(self, name: str) -> None: ...
460460
```
461461

crates/ty_python_semantic/src/types/class.rs

Lines changed: 120 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,9 @@ use crate::semantic_index::{
1919
use crate::types::bound_super::BoundSuperError;
2020
use crate::types::constraints::{ConstraintSet, IteratorConstraintsExtension};
2121
use crate::types::context::InferContext;
22-
use crate::types::diagnostic::{INVALID_TYPE_ALIAS_TYPE, SUPER_CALL_IN_NAMED_TUPLE_METHOD};
22+
use crate::types::diagnostic::{
23+
INVALID_TYPE_ALIAS_TYPE, SUPER_CALL_IN_NAMED_TUPLE_METHOD, UNSOUND_DATACLASS_METHOD_OVERRIDE,
24+
};
2325
use crate::types::enums::enum_metadata;
2426
use crate::types::function::{
2527
DataclassTransformerFlags, DataclassTransformerParams, KnownFunction,
@@ -1936,6 +1938,69 @@ impl<'db> ClassLiteral<'db> {
19361938
Some(typed_dict_params_from_class_def(class_stmt))
19371939
}
19381940

1941+
/// Returns dataclass params for this class, sourced from both dataclass params and dataclass
1942+
/// transform params
1943+
fn merged_dataclass_params(
1944+
self,
1945+
db: &'db dyn Db,
1946+
field_policy: CodeGeneratorKind<'db>,
1947+
) -> (Option<DataclassParams<'db>>, Option<DataclassParams<'db>>) {
1948+
let dataclass_params = self.dataclass_params(db);
1949+
1950+
let mut transformer_params =
1951+
if let CodeGeneratorKind::DataclassLike(Some(transformer_params)) = field_policy {
1952+
Some(DataclassParams::from_transformer_params(
1953+
db,
1954+
transformer_params,
1955+
))
1956+
} else {
1957+
None
1958+
};
1959+
1960+
// Dataclass transformer flags can be overwritten using class arguments.
1961+
if let Some(transformer_params) = transformer_params.as_mut() {
1962+
if let Some(class_def) = self.definition(db).kind(db).as_class() {
1963+
let module = parsed_module(db, self.file(db)).load(db);
1964+
1965+
if let Some(arguments) = &class_def.node(&module).arguments {
1966+
let mut flags = transformer_params.flags(db);
1967+
1968+
for keyword in &arguments.keywords {
1969+
if let Some(arg_name) = &keyword.arg {
1970+
if let Some(is_set) =
1971+
keyword.value.as_boolean_literal_expr().map(|b| b.value)
1972+
{
1973+
for (flag_name, flag) in DATACLASS_FLAGS {
1974+
if arg_name.as_str() == *flag_name {
1975+
flags.set(*flag, is_set);
1976+
}
1977+
}
1978+
}
1979+
}
1980+
}
1981+
1982+
*transformer_params =
1983+
DataclassParams::new(db, flags, transformer_params.field_specifiers(db));
1984+
}
1985+
}
1986+
}
1987+
1988+
(dataclass_params, transformer_params)
1989+
}
1990+
1991+
/// Checks if the given dataclass parameter flag is set for this class.
1992+
/// This checks both the `dataclass_params` and `transformer_params`.
1993+
fn has_dataclass_param(
1994+
self,
1995+
db: &'db dyn Db,
1996+
field_policy: CodeGeneratorKind<'db>,
1997+
param: DataclassFlags,
1998+
) -> bool {
1999+
let (dataclass_params, transformer_params) = self.merged_dataclass_params(db, field_policy);
2000+
dataclass_params.is_some_and(|params| params.flags(db).contains(param))
2001+
|| transformer_params.is_some_and(|params| params.flags(db).contains(param))
2002+
}
2003+
19392004
/// Return the explicit `metaclass` of this class, if one is defined.
19402005
///
19412006
/// ## Note
@@ -2345,57 +2410,8 @@ impl<'db> ClassLiteral<'db> {
23452410
inherited_generic_context: Option<GenericContext<'db>>,
23462411
name: &str,
23472412
) -> Option<Type<'db>> {
2348-
let dataclass_params = self.dataclass_params(db);
2349-
23502413
let field_policy = CodeGeneratorKind::from_class(db, self, specialization)?;
23512414

2352-
let mut transformer_params =
2353-
if let CodeGeneratorKind::DataclassLike(Some(transformer_params)) = field_policy {
2354-
Some(DataclassParams::from_transformer_params(
2355-
db,
2356-
transformer_params,
2357-
))
2358-
} else {
2359-
None
2360-
};
2361-
2362-
// Dataclass transformer flags can be overwritten using class arguments.
2363-
// TODO this should be done more generally, not just in `own_synthesized_member`, so that
2364-
// `dataclass_params` always reflects the transformer params.
2365-
if let Some(transformer_params) = transformer_params.as_mut() {
2366-
if let Some(class_def) = self.definition(db).kind(db).as_class() {
2367-
let module = parsed_module(db, self.file(db)).load(db);
2368-
2369-
if let Some(arguments) = &class_def.node(&module).arguments {
2370-
let mut flags = transformer_params.flags(db);
2371-
2372-
for keyword in &arguments.keywords {
2373-
if let Some(arg_name) = &keyword.arg {
2374-
if let Some(is_set) =
2375-
keyword.value.as_boolean_literal_expr().map(|b| b.value)
2376-
{
2377-
for (flag_name, flag) in DATACLASS_FLAGS {
2378-
if arg_name.as_str() == *flag_name {
2379-
flags.set(*flag, is_set);
2380-
}
2381-
}
2382-
}
2383-
}
2384-
}
2385-
2386-
*transformer_params =
2387-
DataclassParams::new(db, flags, transformer_params.field_specifiers(db));
2388-
}
2389-
}
2390-
}
2391-
2392-
let has_dataclass_param = |param| {
2393-
dataclass_params.is_some_and(|params| params.flags(db).contains(param))
2394-
// TODO if we were correctly initializing `dataclass_params` from the
2395-
// transformer params, this fallback shouldn't be needed here.
2396-
|| transformer_params.is_some_and(|params| params.flags(db).contains(param))
2397-
};
2398-
23992415
let instance_ty =
24002416
Type::instance(db, self.apply_optional_specialization(db, specialization));
24012417

@@ -2514,7 +2530,7 @@ impl<'db> ClassLiteral<'db> {
25142530

25152531
match (field_policy, name) {
25162532
(CodeGeneratorKind::DataclassLike(_), "__init__") => {
2517-
if !has_dataclass_param(DataclassFlags::INIT) {
2533+
if !self.has_dataclass_param(db, field_policy, DataclassFlags::INIT) {
25182534
return None;
25192535
}
25202536

@@ -2529,7 +2545,7 @@ impl<'db> ClassLiteral<'db> {
25292545
signature_from_fields(vec![cls_parameter], Some(Type::none(db)))
25302546
}
25312547
(CodeGeneratorKind::DataclassLike(_), "__lt__" | "__le__" | "__gt__" | "__ge__") => {
2532-
if !has_dataclass_param(DataclassFlags::ORDER) {
2548+
if !self.has_dataclass_param(db, field_policy, DataclassFlags::ORDER) {
25332549
return None;
25342550
}
25352551

@@ -2551,9 +2567,10 @@ impl<'db> ClassLiteral<'db> {
25512567
Some(Type::function_like_callable(db, signature))
25522568
}
25532569
(CodeGeneratorKind::DataclassLike(_), "__hash__") => {
2554-
let unsafe_hash = has_dataclass_param(DataclassFlags::UNSAFE_HASH);
2555-
let frozen = has_dataclass_param(DataclassFlags::FROZEN);
2556-
let eq = has_dataclass_param(DataclassFlags::EQ);
2570+
let unsafe_hash =
2571+
self.has_dataclass_param(db, field_policy, DataclassFlags::UNSAFE_HASH);
2572+
let frozen = self.has_dataclass_param(db, field_policy, DataclassFlags::FROZEN);
2573+
let eq = self.has_dataclass_param(db, field_policy, DataclassFlags::EQ);
25572574

25582575
if unsafe_hash || (frozen && eq) {
25592576
let signature = Signature::new(
@@ -2576,11 +2593,12 @@ impl<'db> ClassLiteral<'db> {
25762593
(CodeGeneratorKind::DataclassLike(_), "__match_args__")
25772594
if Program::get(db).python_version(db) >= PythonVersion::PY310 =>
25782595
{
2579-
if !has_dataclass_param(DataclassFlags::MATCH_ARGS) {
2596+
if !self.has_dataclass_param(db, field_policy, DataclassFlags::MATCH_ARGS) {
25802597
return None;
25812598
}
25822599

2583-
let kw_only_default = has_dataclass_param(DataclassFlags::KW_ONLY);
2600+
let kw_only_default =
2601+
self.has_dataclass_param(db, field_policy, DataclassFlags::KW_ONLY);
25842602

25852603
let fields = self.fields(db, specialization, field_policy);
25862604
let match_args = fields
@@ -2598,8 +2616,8 @@ impl<'db> ClassLiteral<'db> {
25982616
(CodeGeneratorKind::DataclassLike(_), "__weakref__")
25992617
if Program::get(db).python_version(db) >= PythonVersion::PY311 =>
26002618
{
2601-
if !has_dataclass_param(DataclassFlags::WEAKREF_SLOT)
2602-
|| !has_dataclass_param(DataclassFlags::SLOTS)
2619+
if !self.has_dataclass_param(db, field_policy, DataclassFlags::WEAKREF_SLOT)
2620+
|| !self.has_dataclass_param(db, field_policy, DataclassFlags::SLOTS)
26032621
{
26042622
return None;
26052623
}
@@ -2641,7 +2659,7 @@ impl<'db> ClassLiteral<'db> {
26412659
signature_from_fields(vec![self_parameter], Some(instance_ty))
26422660
}
26432661
(CodeGeneratorKind::DataclassLike(_), "__setattr__") => {
2644-
if has_dataclass_param(DataclassFlags::FROZEN) {
2662+
if self.has_dataclass_param(db, field_policy, DataclassFlags::FROZEN) {
26452663
let signature = Signature::new(
26462664
Parameters::new(
26472665
db,
@@ -2662,11 +2680,12 @@ impl<'db> ClassLiteral<'db> {
26622680
(CodeGeneratorKind::DataclassLike(_), "__slots__")
26632681
if Program::get(db).python_version(db) >= PythonVersion::PY310 =>
26642682
{
2665-
has_dataclass_param(DataclassFlags::SLOTS).then(|| {
2666-
let fields = self.fields(db, specialization, field_policy);
2667-
let slots = fields.keys().map(|name| Type::string_literal(db, name));
2668-
Type::heterogeneous_tuple(db, slots)
2669-
})
2683+
self.has_dataclass_param(db, field_policy, DataclassFlags::SLOTS)
2684+
.then(|| {
2685+
let fields = self.fields(db, specialization, field_policy);
2686+
let slots = fields.keys().map(|name| Type::string_literal(db, name));
2687+
Type::heterogeneous_tuple(db, slots)
2688+
})
26702689
}
26712690
(CodeGeneratorKind::TypedDict, "__setitem__") => {
26722691
let fields = self.fields(db, specialization, field_policy);
@@ -3052,6 +3071,42 @@ impl<'db> ClassLiteral<'db> {
30523071
.collect()
30533072
}
30543073

3074+
pub(crate) fn validate_members(self, context: &InferContext<'db, '_>) {
3075+
let db = context.db();
3076+
let Some(field_policy) = CodeGeneratorKind::from_class(db, self, None) else {
3077+
return;
3078+
};
3079+
let class_body_scope = self.body_scope(db);
3080+
let table = place_table(db, class_body_scope);
3081+
let use_def = use_def_map(db, class_body_scope);
3082+
for (symbol_id, declarations) in use_def.all_end_of_scope_symbol_declarations() {
3083+
let result = place_from_declarations(db, declarations.clone());
3084+
let attr = result.ignore_conflicting_declarations();
3085+
let symbol = table.symbol(symbol_id);
3086+
let name = symbol.name();
3087+
if let Some(Type::FunctionLiteral(literal)) = attr.place.ignore_possibly_undefined()
3088+
&& matches!(name.as_str(), "__setattr__" | "__delattr__")
3089+
{
3090+
if let Some(CodeGeneratorKind::DataclassLike(_)) =
3091+
CodeGeneratorKind::from_class(db, self, None)
3092+
&& self.has_dataclass_param(db, field_policy, DataclassFlags::FROZEN)
3093+
{
3094+
if let Some(builder) = context.report_lint(
3095+
&UNSOUND_DATACLASS_METHOD_OVERRIDE,
3096+
literal.node(db, context.file(), context.module()),
3097+
) {
3098+
let mut diagnostic = builder.into_diagnostic(format_args!(
3099+
"Cannot overwrite attribute `{}` in class `{}`",
3100+
name,
3101+
self.name(db)
3102+
));
3103+
diagnostic.info(name);
3104+
}
3105+
}
3106+
}
3107+
}
3108+
}
3109+
30553110
/// Returns a list of all annotated attributes defined in the body of this class. This is similar
30563111
/// to the `__annotations__` attribute at runtime, but also contains default values.
30573112
///

crates/ty_python_semantic/src/types/diagnostic.rs

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,7 @@ pub(crate) fn register_lints(registry: &mut LintRegistryBuilder) {
5050
registry.register_lint(&AMBIGUOUS_PROTOCOL_MEMBER);
5151
registry.register_lint(&CALL_NON_CALLABLE);
5252
registry.register_lint(&POSSIBLY_MISSING_IMPLICIT_CALL);
53+
registry.register_lint(&UNSOUND_DATACLASS_METHOD_OVERRIDE);
5354
registry.register_lint(&CONFLICTING_ARGUMENT_FORMS);
5455
registry.register_lint(&CONFLICTING_DECLARATIONS);
5556
registry.register_lint(&CONFLICTING_METACLASS);
@@ -392,6 +393,32 @@ declare_lint! {
392393
}
393394
}
394395

396+
declare_lint! {
397+
/// ## What it does
398+
/// Checks for dataclass definitions that have both `frozen=True` and a custom `__setattr__` or
399+
/// `__delattr__` method defined.
400+
///
401+
/// ## Why is this bad?
402+
/// Frozen dataclasses synthesize `__setattr__` and `__delattr__` methods which raise a
403+
/// `FrozenInstanceError` to emulate immutability.
404+
///
405+
/// Overriding either of these methods raises a runtime error.
406+
///
407+
/// ## Examples
408+
/// ```python
409+
/// from dataclasses import dataclass
410+
///
411+
/// @dataclass(frozen=True)
412+
/// class A:
413+
/// def __setattr__(self, name: str, value: object) -> None: ...
414+
/// ```
415+
pub(crate) static UNSOUND_DATACLASS_METHOD_OVERRIDE = {
416+
summary: "detects dataclasses with `frozen=True` that have a custom `__setattr__` or `__delattr__` implementation",
417+
status: LintStatus::preview("1.0.0"),
418+
default_level: Level::Error,
419+
}
420+
}
421+
395422
declare_lint! {
396423
/// ## What it does
397424
/// Checks for classes definitions which will fail at runtime due to

crates/ty_python_semantic/src/types/infer/builder.rs

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -990,6 +990,8 @@ impl<'db, 'ast> TypeInferenceBuilder<'db, 'ast> {
990990
if let Some(protocol) = class.into_protocol_class(self.db()) {
991991
protocol.validate_members(&self.context);
992992
}
993+
994+
class.validate_members(&self.context);
993995
}
994996
}
995997

crates/ty_server/tests/e2e/snapshots/e2e__commands__debug_command.snap

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,7 @@ Settings: Settings {
105105
"unresolved-global": Warning (Default),
106106
"unresolved-import": Error (Default),
107107
"unresolved-reference": Error (Default),
108+
"unsound-dataclass-method-override": Error (Default),
108109
"unsupported-base": Warning (Default),
109110
"unsupported-bool-conversion": Error (Default),
110111
"unsupported-operator": Error (Default),

ty.schema.json

Lines changed: 10 additions & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)