Skip to content

Commit 230dd32

Browse files
.
. .
1 parent 86c70f7 commit 230dd32

5 files changed

Lines changed: 49 additions & 48 deletions

File tree

crates/pyrefly_types/src/class.rs

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -249,9 +249,7 @@ impl ClassKind {
249249
("enum", "property") => Self::Property(name.clone()),
250250
("enum", "member") => Self::EnumMember,
251251
("enum", "nonmember") => Self::EnumNonmember,
252-
("dataclasses", "Field")
253-
| ("sqlalchemy.orm", "MappedColumn")
254-
| ("sqlalchemy.orm.properties", "MappedColumn") => Self::DataclassField,
252+
("dataclasses", "Field") => Self::DataclassField,
255253
_ => Self::Class,
256254
}
257255
}

pyrefly/lib/alt/special_calls.rs

Lines changed: 40 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
*/
1313

1414
use pyrefly_types::callable::FuncMetadata;
15+
use pyrefly_types::type_alias::TypeAliasData;
1516
use pyrefly_types::types::Union;
1617
use pyrefly_util::visit::Visit;
1718
use pyrefly_util::visit::VisitMut;
@@ -39,8 +40,6 @@ use crate::types::callable::FunctionKind;
3940
use crate::types::callable::unexpected_keyword;
4041
use crate::types::class::Class;
4142
use crate::types::class::ClassType;
42-
use crate::types::special_form::SpecialForm;
43-
use crate::types::special_form::SpecialForm;
4443
use crate::types::tuple::Tuple;
4544
use crate::types::types::Type;
4645

@@ -338,7 +337,7 @@ impl<'a, Ans: LookupAnswer> AnswersSolver<'a, Ans> {
338337
}
339338
for keyword in &call.arguments.keywords {
340339
if let Some(arg) = &keyword.arg
341-
&& matches!(arg.as_str(), "__type_pos" | "type_")
340+
&& matches!(arg.as_str(), "type_")
342341
&& let Some(ty) = self.python_type_from_type_engine_expr(&keyword.value)
343342
{
344343
return Some(ty);
@@ -368,7 +367,7 @@ impl<'a, Ans: LookupAnswer> AnswersSolver<'a, Ans> {
368367
_ => {}
369368
}
370369
}
371-
if primary_key { false } else { nullable }
370+
!primary_key && nullable
372371
}
373372

374373
fn expr_bool_literal(expr: &Expr) -> Option<bool> {
@@ -391,61 +390,70 @@ impl<'a, Ans: LookupAnswer> AnswersSolver<'a, Ans> {
391390
self.python_type_from_type_engine_type(&inst)
392391
}
393392
Type::Type(inner) => self.python_type_from_type_engine_type(inner),
394-
Type::TypeAlias(alias) => self.python_type_from_type_engine_type(&alias.as_type()),
393+
Type::TypeAlias(alias) => match alias.as_ref() {
394+
TypeAliasData::Value(alias) => {
395+
self.python_type_from_type_engine_type(&alias.as_type())
396+
}
397+
TypeAliasData::Ref(_) => {
398+
self.python_type_from_type_engine_type(&self.untype_alias(alias))
399+
}
400+
},
395401
Type::Union(u) => {
396-
let mut inferred = None;
402+
let mut inferred = Vec::new();
397403
for member in &u.members {
398404
if let Some(member_ty) = self.python_type_from_type_engine_type(member) {
399-
match &inferred {
400-
Some(existing) if existing != &member_ty => return None,
401-
None => inferred = Some(member_ty),
402-
_ => {}
403-
}
405+
inferred.push(member_ty);
406+
}
407+
}
408+
if inferred.is_empty() {
409+
None
410+
} else {
411+
inferred.sort();
412+
inferred.dedup();
413+
if inferred.len() == 1 {
414+
inferred.pop()
415+
} else {
416+
Some(Type::union(inferred))
404417
}
405418
}
406-
inferred
407419
}
408420
_ => None,
409421
}
410422
}
411423

412424
fn python_type_from_type_engine_class(&self, cls: &ClassType) -> Option<Type> {
413-
if Self::is_sqlalchemy_type_engine_class(cls.class_object()) {
425+
if cls
426+
.class_object()
427+
.has_toplevel_qname("sqlalchemy.sql.type_api", "TypeEngine")
428+
{
414429
return cls.targs().as_slice().first().cloned();
415430
}
416-
let bases = self.get_base_types_for_class(cls.class_object());
417-
for base in bases.iter() {
418-
if let Some(ty) = self.python_type_from_type_engine_class(base) {
419-
return Some(ty);
431+
let mro = self.get_mro_for_class(cls.class_object());
432+
for ancestor in mro.ancestors_no_object() {
433+
if ancestor
434+
.class_object()
435+
.has_toplevel_qname("sqlalchemy.sql.type_api", "TypeEngine")
436+
{
437+
return ancestor.targs().as_slice().first().cloned();
420438
}
421439
}
422440
None
423441
}
424442

425-
fn is_sqlalchemy_type_engine_class(class: &Class) -> bool {
426-
class.has_toplevel_qname("sqlalchemy.sql.type_api", "TypeEngine")
427-
}
428-
429443
fn apply_sqlalchemy_mapped_python_type(&self, mut ty: Type, python_type: Type) -> Type {
430444
ty.visit_mut(&mut |inner| {
431445
if let Type::ClassType(class_type) = inner
432-
&& Self::is_sqlalchemy_mapped_class(class_type.class_object())
446+
&& class_type
447+
.class_object()
448+
.has_toplevel_qname("sqlalchemy.orm.properties", "MappedColumn")
449+
&& let Some(slot) = class_type.targs_mut().as_mut().get_mut(0)
433450
{
434-
if let Some(slot) = class_type.targs_mut().as_mut().get_mut(0) {
435-
*slot = python_type.clone();
436-
}
451+
*slot = python_type.clone();
437452
}
438453
});
439454
ty
440455
}
441456

442-
fn is_sqlalchemy_mapped_class(class: &Class) -> bool {
443-
class.has_toplevel_qname("sqlalchemy.orm.base", "Mapped")
444-
|| class.has_toplevel_qname("sqlalchemy.orm.base", "_MappedAnnotationBase")
445-
|| class.has_toplevel_qname("sqlalchemy.orm.base", "_DeclarativeMapped")
446-
|| class.has_toplevel_qname("sqlalchemy.orm.properties", "MappedColumn")
447-
}
448-
449457
pub fn call_isinstance(
450458
&self,
451459
obj: &Expr,

pyrefly/lib/export/special.rs

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -205,13 +205,8 @@ impl SpecialExport {
205205
m.as_str(),
206206
"typing" | "typing_extensions" | "collections.abc"
207207
),
208+
Self::SqlAlchemyMappedColumn => matches!(m.as_str(), "sqlalchemy.orm"),
208209
Self::Deprecated => matches!(m.as_str(), "warnings" | "typing_extensions"),
209-
Self::SqlAlchemyMappedColumn => {
210-
matches!(
211-
m.as_str(),
212-
"sqlalchemy.orm" | "sqlalchemy.orm._orm_constructors"
213-
)
214-
}
215210
}
216211
}
217212

pyrefly/lib/test/mod.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -65,8 +65,8 @@ mod returns;
6565
mod scope;
6666
mod semantic_syntax_errors;
6767
mod simple;
68-
mod sqlalchemy;
6968
mod slots;
69+
mod sqlalchemy;
7070
mod state;
7171
mod subscript_narrow;
7272
mod suppression;

pyrefly/lib/test/sqlalchemy.rs

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -102,8 +102,8 @@ testcase!(
102102
test_sqlalchemy_mapped_column_infers_type,
103103
sqlalchemy_env(),
104104
r#"
105-
from typing import reveal_type
106-
from sqlalchemy.orm import mapped_column
105+
from typing import assert_type
106+
from sqlalchemy.orm import MappedColumn, mapped_column
107107
from sqlalchemy.sql.sqltypes import Integer, String
108108
109109
class Model:
@@ -112,9 +112,9 @@ class Model:
112112
sku = mapped_column(String(), nullable=False)
113113
pk = mapped_column(Integer, primary_key=True)
114114
115-
reveal_type(Model.name) # E: revealed type: MappedColumn[str | None]
116-
reveal_type(Model.quantity) # E: revealed type: MappedColumn[int | None]
117-
reveal_type(Model.sku) # E: revealed type: MappedColumn[str]
118-
reveal_type(Model.pk) # E: revealed type: MappedColumn[int]
115+
assert_type(Model.name, MappedColumn[str | None])
116+
assert_type(Model.quantity, MappedColumn[int | None])
117+
assert_type(Model.sku, MappedColumn[str])
118+
assert_type(Model.pk, MappedColumn[int])
119119
"#,
120120
);

0 commit comments

Comments
 (0)