Files
kjol/tools/tsgo/internal/ls/autoimport/import_adder.go
2026-07-09 16:50:43 -04:00

502 lines
18 KiB
Go

package autoimport
import (
"context"
"fmt"
"maps"
"slices"
"github.com/microsoft/typescript-go/internal/ast"
"github.com/microsoft/typescript-go/internal/checker"
"github.com/microsoft/typescript-go/internal/compiler"
"github.com/microsoft/typescript-go/internal/core"
"github.com/microsoft/typescript-go/internal/debug"
"github.com/microsoft/typescript-go/internal/locale"
// "github.com/microsoft/typescript-go/internal/ls"
"github.com/microsoft/typescript-go/internal/ls/change"
"github.com/microsoft/typescript-go/internal/ls/lsconv"
"github.com/microsoft/typescript-go/internal/ls/lsutil"
"github.com/microsoft/typescript-go/internal/lsp/lsproto"
"github.com/microsoft/typescript-go/internal/nodebuilder"
)
type ImportAdder interface {
HasFixes() bool
AddImportFromExportedSymbol(symbol *ast.Symbol, isValidTypeOnlyUseSite bool)
AddImportFix(fix *Fix)
Edits() []*lsproto.TextEdit
}
// addToExistingState tracks modifications to an existing import clause or binding pattern
type addToExistingState struct {
importClauseOrBindingPattern *ast.ImportClauseOrBindingPattern
defaultImport *newImportBinding
namedImports map[string]*newImportBinding
}
// importsCollection tracks new imports to be created for a given module specifier
type importsCollection struct {
defaultImport *newImportBinding
namedImports map[string]*newImportBinding
namespaceLikeImport *newImportBinding
useRequire bool
}
func newImportsKey(moduleSpecifier string, topLevelTypeOnly bool) string {
if topLevelTypeOnly {
return "1|" + moduleSpecifier
}
return "0|" + moduleSpecifier
}
type importAdder struct {
// Context
ctx context.Context
checker *checker.Checker
view *View
formatOptions lsutil.FormatCodeSettings
converters *lsconv.Converters
preferences lsutil.UserPreferences
// State
addToNamespace []*Fix // Namespace fixes don't conflict, so just build a list
importType []*Fix // JSDoc type import fixes
addToExisting map[*ast.ImportClauseOrBindingPattern]*addToExistingState // importClauseOrBindingPattern -> default or named bindings
newImports map[string]*importsCollection // module specifier + type only -> imports
// !!! removeExisting, verbatimImports?
}
func NewImportAdder(
ctx context.Context,
program *compiler.Program,
checker *checker.Checker,
file *ast.SourceFile,
view *View,
formatOptions lsutil.FormatCodeSettings,
converters *lsconv.Converters,
preferences lsutil.UserPreferences,
) ImportAdder {
return &importAdder{
ctx: ctx,
checker: checker,
view: view,
formatOptions: formatOptions,
converters: converters,
preferences: preferences,
addToNamespace: nil,
importType: nil,
addToExisting: make(map[*ast.Node]*addToExistingState),
newImports: make(map[string]*importsCollection),
}
}
func (adder *importAdder) HasFixes() bool {
return len(adder.addToNamespace) > 0 ||
len(adder.importType) > 0 ||
len(adder.addToExisting) > 0 ||
len(adder.newImports) > 0
}
// !!! referenceImport
func (adder *importAdder) AddImportFromExportedSymbol(exportedSymbol *ast.Symbol, isValidTypeOnlyUseSite bool) {
symbol := adder.checker.GetMergedSymbol(adder.checker.SkipAlias(exportedSymbol))
exportInfos := adder.getAllExportsForSymbol(symbol)
if len(exportInfos) == 0 {
// If no exportInfo is found, this means export could not be resolved when we have filtered for autoImportFileExcludePatterns,
// so we should not generate an import.
// debug.Assert(len(adder.ls.UserPreferences().AutoImportFileExcludePatterns) > 0)
return
}
fix := adder.getImportFixForSymbol(adder.view, adder.view.importingFile, exportInfos, isValidTypeOnlyUseSite)
if fix != nil {
// !!! referenceImport -> propertyName
adder.AddImportFix(fix)
}
}
func (adder *importAdder) Edits() []*lsproto.TextEdit {
// !!! organize imports?
tracker := change.NewTracker(adder.ctx, adder.view.program.Options(), adder.formatOptions, adder.converters)
quotePreference := lsutil.GetQuotePreference(adder.view.importingFile, adder.preferences)
for _, fix := range adder.addToNamespace {
addNamespaceQualifier(fix, tracker, adder.view.importingFile, locale.Default)
}
for _, fix := range adder.importType {
addImportType(fix, adder.view.importingFile, adder.preferences, tracker, locale.Default)
}
for clauseOrPattern, entry := range adder.addToExisting {
addToExistingImport(
tracker,
adder.view.importingFile,
clauseOrPattern,
entry.defaultImport,
sortedNamedImports(entry.namedImports),
adder.preferences,
)
}
var newDeclarations []*ast.AnyImportOrRequireStatement
for key, newImport := range adder.newImports {
moduleSpecifier := key[2:] // From `${0 | 1}|${moduleSpecifier}` format
var declarations []*ast.AnyImportOrRequireStatement
if newImport.useRequire {
declarations = getNewRequires(
tracker,
moduleSpecifier,
quotePreference,
newImport.defaultImport,
sortedNamedImports(newImport.namedImports),
newImport.namespaceLikeImport,
adder.view.program.Options(),
)
} else {
declarations = getNewImports(
tracker,
moduleSpecifier,
quotePreference,
newImport.defaultImport,
sortedNamedImports(newImport.namedImports),
newImport.namespaceLikeImport,
adder.view.program.Options(),
adder.preferences,
)
}
newDeclarations = append(newDeclarations, declarations...)
}
if len(newDeclarations) > 0 {
insertImports(tracker, adder.view.importingFile, newDeclarations, true /*blankLineBetween*/, adder.preferences)
}
return tracker.GetChanges()[adder.view.importingFile.FileName()]
}
func sortedNamedImports(m map[string]*newImportBinding) []*newImportBinding {
keys := slices.Sorted(maps.Keys(m))
result := make([]*newImportBinding, 0, len(keys))
for _, k := range keys {
result = append(result, m[k])
}
return result
}
// AddImportFix adds a fix to the import adder, accumulating it with other fixes
// so that multiple imports from the same module are coalesced into a single import statement.
func (adder *importAdder) AddImportFix(fix *Fix) {
symbolName := fix.Name
compilerOptions := adder.view.program.Options()
switch fix.Kind {
case lsproto.AutoImportFixKindUseNamespace:
adder.addToNamespace = append(adder.addToNamespace, fix)
case lsproto.AutoImportFixKindJsdocTypeImport:
adder.importType = append(adder.importType, fix)
case lsproto.AutoImportFixKindAddToExisting:
existingFix := getAddToExistingImportFix(adder.view.importingFile, fix)
entry := adder.addToExisting[existingFix.importClauseOrBindingPattern]
if entry == nil {
entry = &addToExistingState{
importClauseOrBindingPattern: existingFix.importClauseOrBindingPattern,
namedImports: make(map[string]*newImportBinding),
}
adder.addToExisting[existingFix.importClauseOrBindingPattern] = entry
}
if fix.ImportKind == lsproto.ImportKindNamed {
prevImport := entry.namedImports[symbolName]
var prevTypeOnly lsproto.AddAsTypeOnly
if prevImport != nil {
prevTypeOnly = prevImport.addAsTypeOnly
}
entry.namedImports[symbolName] = &newImportBinding{
kind: lsproto.ImportKindNamed,
name: symbolName,
addAsTypeOnly: reduceAddAsTypeOnlyValues(prevTypeOnly, fix.AddAsTypeOnly),
propertyName: existingFix.namedImport.propertyName,
}
} else {
// Default import
debug.Assert(
entry.defaultImport == nil || entry.defaultImport.name == symbolName,
"(Add to Existing) Default import should be missing or match symbolName",
)
var prevTypeOnly lsproto.AddAsTypeOnly
if entry.defaultImport != nil {
prevTypeOnly = entry.defaultImport.addAsTypeOnly
}
entry.defaultImport = &newImportBinding{
kind: lsproto.ImportKindDefault,
name: symbolName,
addAsTypeOnly: reduceAddAsTypeOnlyValues(prevTypeOnly, fix.AddAsTypeOnly),
}
}
case lsproto.AutoImportFixKindAddNew:
entry := adder.getNewImportEntry(fix.ModuleSpecifier, fix.ImportKind, fix.UseRequire, fix.AddAsTypeOnly)
debug.Assert(
entry.useRequire == fix.UseRequire,
"(Add new) Tried to add an `import` and a `require` for the same module",
)
switch fix.ImportKind {
case lsproto.ImportKindDefault:
debug.Assert(
entry.defaultImport == nil || entry.defaultImport.name == symbolName,
"(Add new) Default import should be missing or match symbolName",
)
var prevTypeOnly lsproto.AddAsTypeOnly
if entry.defaultImport != nil {
prevTypeOnly = entry.defaultImport.addAsTypeOnly
}
entry.defaultImport = &newImportBinding{
kind: lsproto.ImportKindDefault,
name: symbolName,
addAsTypeOnly: reduceAddAsTypeOnlyValues(prevTypeOnly, fix.AddAsTypeOnly),
}
case lsproto.ImportKindNamed:
if entry.namedImports == nil {
entry.namedImports = make(map[string]*newImportBinding)
}
prevImport := entry.namedImports[symbolName]
var prevTypeOnly lsproto.AddAsTypeOnly
if prevImport != nil {
prevTypeOnly = prevImport.addAsTypeOnly
}
entry.namedImports[symbolName] = &newImportBinding{
kind: lsproto.ImportKindNamed,
name: symbolName,
addAsTypeOnly: reduceAddAsTypeOnlyValues(prevTypeOnly, fix.AddAsTypeOnly),
// !!! propertyName
}
case lsproto.ImportKindCommonJS:
if compilerOptions.VerbatimModuleSyntax == core.TSTrue {
if entry.namedImports == nil {
entry.namedImports = make(map[string]*newImportBinding)
}
prevImport := entry.namedImports[symbolName]
var prevTypeOnly lsproto.AddAsTypeOnly
if prevImport != nil {
prevTypeOnly = prevImport.addAsTypeOnly
}
entry.namedImports[symbolName] = &newImportBinding{
kind: lsproto.ImportKindCommonJS,
name: symbolName,
addAsTypeOnly: reduceAddAsTypeOnlyValues(prevTypeOnly, fix.AddAsTypeOnly),
// !!! propertyName
}
} else {
debug.Assert(
entry.namespaceLikeImport == nil || entry.namespaceLikeImport.name == symbolName,
"Namespacelike import should be missing or match symbolName",
)
entry.namespaceLikeImport = &newImportBinding{
kind: lsproto.ImportKindCommonJS,
name: symbolName,
addAsTypeOnly: fix.AddAsTypeOnly,
}
}
case lsproto.ImportKindNamespace:
debug.Assert(
entry.namespaceLikeImport == nil || entry.namespaceLikeImport.name == symbolName,
"Namespacelike import should be missing or match symbolName",
)
entry.namespaceLikeImport = &newImportBinding{
kind: lsproto.ImportKindNamespace,
name: symbolName,
addAsTypeOnly: fix.AddAsTypeOnly,
}
}
case lsproto.AutoImportFixKindPromoteTypeOnly:
// Excluding from fix-all
default:
debug.Fail(fmt.Sprintf("Unexpected fix kind: %v", fix.Kind))
}
}
// `NotAllowed` overrides `Required` because one addition of a new import might be required to be type-only
// because of `--importsNotUsedAsValues=error`, but if a second addition of the same import is `NotAllowed`
// to be type-only, the reason the first one was `Required` - the unused runtime dependency - is now moot.
// Alternatively, if one addition is `Required` because it has no value meaning under `--preserveValueImports`
// and `--isolatedModules`, it should be impossible for another addition to be `NotAllowed` since that would
// mean a type is being referenced in a value location.
func reduceAddAsTypeOnlyValues(prevValue, newValue lsproto.AddAsTypeOnly) lsproto.AddAsTypeOnly {
if newValue > prevValue {
return newValue
}
return prevValue
}
func (adder *importAdder) getNewImportEntry(moduleSpecifier string, importKind lsproto.ImportKind, useRequire bool, addAsTypeOnly lsproto.AddAsTypeOnly) *importsCollection {
// A default import that requires type-only makes the whole import type-only.
// (We could add `default` as a named import, but that style seems undesirable.)
// Under `--preserveValueImports` and `--importsNotUsedAsValues=error`, if a
// module default-exports a type but named-exports some values (weird), you would
// have to use a type-only default import and non-type-only named imports. These
// require two separate import declarations, so we build this into the map key.
typeOnlyKey := newImportsKey(moduleSpecifier, true /*topLevelTypeOnly*/)
nonTypeOnlyKey := newImportsKey(moduleSpecifier, false /*topLevelTypeOnly*/)
typeOnlyEntry := adder.newImports[typeOnlyKey]
nonTypeOnlyEntry := adder.newImports[nonTypeOnlyKey]
newEntry := &importsCollection{
useRequire: useRequire,
}
if importKind == lsproto.ImportKindDefault && addAsTypeOnly == lsproto.AddAsTypeOnlyRequired {
if typeOnlyEntry != nil {
return typeOnlyEntry
}
adder.newImports[typeOnlyKey] = newEntry
return newEntry
}
if addAsTypeOnly == lsproto.AddAsTypeOnlyAllowed && (typeOnlyEntry != nil || nonTypeOnlyEntry != nil) {
if typeOnlyEntry != nil {
return typeOnlyEntry
}
return nonTypeOnlyEntry
}
if nonTypeOnlyEntry != nil {
return nonTypeOnlyEntry
}
adder.newImports[nonTypeOnlyKey] = newEntry
return newEntry
}
func (adder *importAdder) getAllExportsForSymbol(
symbol *ast.Symbol,
) []*Export {
if export := SymbolToExport(symbol, adder.checker); export != nil {
return adder.view.SearchByExportID(export.ExportID)
}
return nil
}
func TypeToAutoImportableTypeNode(
c *checker.Checker,
importAdder ImportAdder,
t *checker.Type,
contextNode *ast.Node, // !!! flags
) *ast.TypeNode {
idToSymbol := make(map[*ast.IdentifierNode]*ast.Symbol)
typeNode := c.TypeToTypeNode(t, contextNode, nodebuilder.FlagsNone, idToSymbol)
if typeNode == nil {
return nil
}
return TypeNodeToAutoImportableTypeNode(typeNode, importAdder, idToSymbol)
}
// TypeNodeToAutoImportableTypeNode converts import type references in a type node to
// simple type references and registers needed imports with the import adder.
func TypeNodeToAutoImportableTypeNode(
typeNode *ast.TypeNode,
importAdder ImportAdder,
idToSymbol map[*ast.IdentifierNode]*ast.Symbol,
) *ast.TypeNode {
referenceTypeNode, importableSymbols := TryGetAutoImportableReferenceFromTypeNode(typeNode, idToSymbol)
if referenceTypeNode != nil {
if importAdder != nil {
importSymbols(importAdder, importableSymbols)
}
typeNode = referenceTypeNode
}
// !!! handle type node reuse: nodes needs to be fresh here but also preserve symbols
return typeNode
}
func importSymbols(importAdder ImportAdder, symbols []*ast.Symbol) {
for _, symbol := range symbols {
importAdder.AddImportFromExportedSymbol(symbol, true /*isValidTypeOnlyUseSite*/)
}
}
// Given a type node containing 'import("./a").SomeType<import("./b").OtherType<...>>',
// returns an equivalent type reference node with any nested ImportTypeNodes also replaced
// with type references, and a list of symbols that must be imported to use the type reference.
// TryGetAutoImportableReferenceFromTypeNode converts import type references in a type node
// to simple type references and returns the transformed type node and the symbols that need
// to be imported.
func TryGetAutoImportableReferenceFromTypeNode(importTypeNode *ast.TypeNode, idToSymbol map[*ast.IdentifierNode]*ast.Symbol) (*ast.TypeNode, []*ast.Symbol) {
var symbols []*ast.Symbol
var visitor *ast.NodeVisitor
factory := ast.NewNodeFactory(ast.NodeFactoryHooks{})
visit := func(node *ast.Node) *ast.Node {
if ast.IsLiteralImportTypeNode(node) && node.AsImportTypeNode().Qualifier != nil {
importTypeNode := node.AsImportTypeNode()
// Symbol for the left-most thing after the dot
firstIdentifier := ast.GetFirstIdentifier(importTypeNode.Qualifier)
symbol := idToSymbol[firstIdentifier]
if symbol == nil {
// if symbol is missing then this doesn't come from a synthesized import type node
// it has to be an import type node authored by the user and thus it has to be valid
// it can't refer to reserved internal symbol names and such
return node.VisitEachChild(visitor)
}
name := getNameForExportedSymbol(symbol, false /*preferCapitalized*/)
var qualifier *ast.EntityName
if name != firstIdentifier.Text() {
qualifier = replaceFirstIdentifierOfEntityName(factory, importTypeNode.Qualifier, factory.NewIdentifier(name))
} else {
qualifier = importTypeNode.Qualifier
}
symbols = append(symbols, symbol)
typeArguments := visitor.VisitNodes(importTypeNode.TypeArguments)
return factory.NewTypeReferenceNode(qualifier, typeArguments)
}
return visitor.VisitEachChild(node)
}
visitor = ast.NewNodeVisitor(visit, factory, ast.NodeVisitorHooks{})
typeNode := visitor.VisitNode(importTypeNode)
debug.Assert(typeNode == nil || ast.IsTypeNode(typeNode), "expected a type node")
return typeNode, symbols
}
// If a type checker and multiple files are available, consider using `forEachNameOfDefaultExport`
// instead, which searches for names of re-exported defaults/namespaces in target files.
func getNameForExportedSymbol(symbol *ast.Symbol, preferCapitalized bool) string {
if symbol.Name == ast.InternalSymbolNameExportEquals || symbol.Name == ast.InternalSymbolNameDefault {
// Names for default exports:
// - export default foo => foo
// - export { foo as default } => foo
// - export default 0 => filename converted to camelCase
name := getDefaultLikeExportNameFromDeclaration(symbol)
if name != "" {
return name
}
debug.Assert(symbol.Parent != nil, "Expected exported symbol to have module symbol as parent")
return lsutil.ModuleSymbolToValidIdentifier(symbol.Parent, preferCapitalized)
}
return symbol.Name
}
func replaceFirstIdentifierOfEntityName(factory *ast.NodeFactory, name *ast.EntityName, newIdentifier *ast.IdentifierNode) *ast.EntityName {
if name.Kind == ast.KindIdentifier {
return newIdentifier
}
return factory.NewQualifiedName(
replaceFirstIdentifierOfEntityName(factory, name.AsQualifiedName().Left, newIdentifier),
name.AsQualifiedName().Right,
)
}
func (adder *importAdder) getImportFixForSymbol(view *View, file *ast.SourceFile, exports []*Export, isValidTypeOnlyUseSite bool) *Fix {
fixes := core.FlatMap(exports, func(export *Export) []*Fix {
return view.GetFixes(adder.ctx, export, false /*forJSX*/, isValidTypeOnlyUseSite, nil /*usagePosition*/)
})
slices.SortFunc(fixes, func(a, b *Fix) int {
return view.CompareFixesForRanking(a, b)
})
if len(fixes) > 0 {
return fixes[0]
}
return nil
}